Polish Lamb Kernel (#7299)

This commit is contained in:
Weixing Zhang 2021-04-09 09:55:57 -07:00 committed by GitHub
parent 711cc99f4d
commit c22963c23d
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 76 additions and 84 deletions

View file

@ -230,13 +230,13 @@ Status launch_lamb_compute_direction(
p_m2s[i],
p_loss_scale,
p_g_norm,
CudaT4(alphas[i]),
CudaT4(betas[i]),
CudaT2(lambdas[i]),
CudaT4(epsilons[i]),
CudaT2(max_norms[i]),
CudaT4(alpha_correction),
CudaT4(beta_correction),
alphas[i],
betas[i],
lambdas[i],
epsilons[i],
max_norms[i],
alpha_correction,
beta_correction,
p_ds[i],
p_m1_news[i],
p_m2_news[i],
@ -276,7 +276,7 @@ Status launch_lamb_compute_direction(
tensor_sizes_in_buckets[key],
buckets[key],
lamb_stage1,
p_loss_scale, p_g_norm, lambda, alpha, beta, epsilon, CudaT2(max_norm), alpha_correction, beta_correction);
p_loss_scale, p_g_norm, lambda, alpha, beta, epsilon, max_norm, alpha_correction, beta_correction);
}
return Status::OK();

View file

@ -20,41 +20,40 @@ __device__ __forceinline__ void _LambComputeDirectionRule(
const T2& g,
const T3& m1,
const T3& m2,
const T3& alpha,
const T3& beta,
const T1& lambda,
const T3& epsilon,
const T3& alpha_correction,
const T3& beta_correction,
const float& alpha,
const float& beta,
const float& lambda,
const float& epsilon,
const float& alpha_correction,
const float& beta_correction,
T2& d,
T3& m1_new,
T3& m2_new) {
// Actual gradient. The scale is a product of loss' scale and
// global gradient norm (if the norm > 1).
const T3 g_unscaled = T3(T1(g) / g_scale);
const T1 g_unscaled = T1(g) / g_scale;
// A constant in Lamb's equation.
const T3 one = T3(1.0f);
const T1 one = T1(1.0f);
// Update exponentially-averaged historical gradient
const T3 m1_new_tmp = alpha * m1 + (one - alpha) * g_unscaled;
const T1 m1_new_tmp = alpha * static_cast<T1>(m1) + (one - alpha) * g_unscaled;
// Update exponentially-averaged historical squared gradient
const T3 m2_new_tmp = beta * m2 + (one - beta) * g_unscaled * g_unscaled;
const T1 m2_new_tmp = beta * static_cast<T1>(m2) + (one - beta) * g_unscaled * g_unscaled;
// Compute unbiased 1st-order momentom.
// The value alpha_correction is usually (1-alpha^t),
// where t is the number of executed training iterations.
const T3 m1_new_tmp_corrected = m1_new_tmp / alpha_correction;
const T1 m1_new_tmp_corrected = m1_new_tmp / alpha_correction;
// Compute unbiased 2nd-order momentom.
// The value beta_correction is usually (1-beta^t),
// where t is the number of executed training iterations.
const T3 m2_new_tmp_corrected = m2_new_tmp / beta_correction;
const T1 m2_new_tmp_corrected = m2_new_tmp / beta_correction;
// Save regularized update direction to output.
const T2 d_tmp = lambda * w +
T1(m1_new_tmp_corrected / (_Sqrt(m2_new_tmp_corrected) + epsilon));
const T1 d_tmp = lambda * w + m1_new_tmp_corrected / (_Sqrt(m2_new_tmp_corrected) + epsilon);
// Things are updated only if the direction is finite.
if (_IsFiniteScalar(d_tmp)) {
@ -76,13 +75,13 @@ __global__ void _LambComputeDirectionImpl(
const T3* moment_2,
const T1* loss_scale,
const T_GRAD_NORM* g_norm,
T3 alpha,
T3 beta,
T1 lambda,
T3 epsilon,
T1 max_norm,
T3 alpha_correction,
T3 beta_correction,
float alpha,
float beta,
float lambda,
float epsilon,
float max_norm,
float alpha_correction,
float beta_correction,
T2* update_direction,
T3* moment_1_out,
T3* moment_2_out,
@ -117,13 +116,13 @@ void LambComputeDirection(
const T3* moment_2,
const T1* loss_scale,
const T_GRAD_NORM* grad_norm,
T3 alpha,
T3 beta,
T1 lambda,
T3 epsilon,
T1 max_norm,
T3 alpha_correction,
T3 beta_correction,
float alpha,
float beta,
float lambda,
float epsilon,
float max_norm,
float alpha_correction,
float beta_correction,
T2* update_direction,
T3* moment_1_out,
T3* moment_2_out,
@ -160,13 +159,13 @@ void LambComputeDirection(
const T3* moment_2, \
const T1* loss_scale, \
const T_GRAD_NORM* grad_norm, \
T3 alpha, \
T3 beta, \
T1 lambda, \
T3 epsilon, \
T1 max_norm, \
T3 alpha_correction, \
T3 beta_correction, \
float alpha, \
float beta, \
float lambda, \
float epsilon, \
float max_norm, \
float alpha_correction, \
float beta_correction, \
T2* weights_out, \
T3* moment_1_out, \
T3* moment_2_out, \
@ -319,13 +318,13 @@ __global__ void LambMultiTensorComputeDirectionImpl(
ChunkGroup<6> chunk_group,
const T1* loss_scale,
const T_GRAD_NORM* g_norm,
const T1 lambda,
const T3 alpha,
const T3 beta,
const T3 epsilon,
const float lambda,
const float alpha,
const float beta,
const float epsilon,
const T1 max_norm,
const T3 alpha_correction,
const T3 beta_correction) {
const float alpha_correction,
const float beta_correction) {
const int group_index = chunk_group.block_index_to_tensor_group_index[blockIdx.x];
const int tensor_size = chunk_group.tensor_sizes[group_index];
const int chunk_size = chunk_group.chunk_size;
@ -364,13 +363,13 @@ void LambMultiTensorComputeDirectionFunctor<T1, T2, T3, T_GRAD_NORM>::operator()
ChunkGroup<6> chunk_group,
const T1* loss_scale,
const T_GRAD_NORM* g_norm,
const T1 lambda,
const T3 alpha,
const T3 beta,
const T3 epsilon,
const T1 max_norm,
const T3 alpha_correction,
const T3 beta_correction) {
const float lambda,
const float alpha,
const float beta,
const float epsilon,
const float max_norm,
const float alpha_correction,
const float beta_correction) {
const int thread_count = ChunkGroup<6>::thread_count_per_block;
const int block_count = chunk_group.chunk_count;
@ -393,13 +392,13 @@ void LambMultiTensorComputeDirectionFunctor<T1, T2, T3, T_GRAD_NORM>::operator()
ChunkGroup<6> chunk_group, \
const T1* loss_scale, \
const T_GRAD_NORM* g_norm, \
const T1 lambda, \
const T3 alpha, \
const T3 beta, \
const T3 epsilon, \
const T1 max_norm, \
const T3 alpha_correction, \
const T3 beta_correction);
const float lambda, \
const float alpha, \
const float beta, \
const float epsilon, \
const float max_norm, \
const float alpha_correction, \
const float beta_correction);
INSTANTIATE_LAMB_STAGE1_MULTI_TENSOR_FUNCTOR(float, float, float, float)
INSTANTIATE_LAMB_STAGE1_MULTI_TENSOR_FUNCTOR(double, double, double, double)

View file

@ -56,13 +56,13 @@ void LambComputeDirection(
const T3* moment_2,
const T1* loss_scale,
const T_GRAD_NORM* grad_norm,
T3 alpha,
T3 beta,
T1 lambda,
T3 epsilon,
T1 max_norm,
T3 alpha_correction,
T3 beta_correction,
float alpha,
float beta,
float lambda,
float epsilon,
float max_norm,
float alpha_correction,
float beta_correction,
T2* update_direction,
T3* moment_1_out,
T3* moment_2_out,
@ -112,13 +112,13 @@ struct LambMultiTensorComputeDirectionFunctor {
ChunkGroup<6> chunk_group,
const T1* loss_scale,
const T_GRAD_NORM* grad_norm,
const T1 lambda,
const T3 alpha,
const T3 beta,
const T3 epsilon,
const T1 max_norm,
const T3 alpha_correction,
const T3 beta_correction);
const float lambda,
const float alpha,
const float beta,
const float epsilon,
const float max_norm,
const float alpha_correction,
const float beta_correction);
};
// Lamb's reduction maps [w, d] to [w_norm, d_norm] where

View file

@ -1,10 +1,3 @@
OptimizerTest.LambOptimizerTestExternalBaselineDouble
OptimizerTest.LambOptimizerTest5DTensorMixPrecision32_16
OptimizerTest.LambOptimizerTestSimpleBaselineMixPrecision32_16
OptimizerTest.LambOptimizerTestBaselineMixPrecision32_16
OptimizerTest.LambOptimizerTestScalarMixPrecision32_16
OptimizerTest.LambOptimizerTestScalarMixPrecision32_16_NoDefaultMaxNormClipping
OptimizerTest.LambOptimizerTestLarge
CudaKernelTest.SparseSoftmaxCrossEntropy_LargeSizeTensor
CudaKernelTest.NegativeLogLikelihoodLoss_TinySizeTensor
CudaKernelTest.NegativeLogLikelihoodLoss_SmallSizeTensor