Ë
    •\;j�2  ã                   ó^   — d dl mZ d dlmZ ddlmZmZ ddlmZ ddl	m
Z
 g Z G d„ d	e
«      Zy
)é    )Ú_C_ops)Úglobal_scopeé   )ÚcoreÚ	framework)ÚVariableé   )Ú	Optimizerc                   ój   ‡ — e Zd ZdZdZdZdZdZ	 	 	 	 	 	 	 	 	 	 	 dˆ fd„	Zdd„Z	d„ Z
d	„ Zd
„ Zd„ Zˆ xZS )ÚLambab  
    LAMB (Layer-wise Adaptive Moments optimizer for Batching training) Optimizer.

    LAMB Optimizer is designed to scale up the batch size of training without losing
    accuracy, which supports adaptive element-wise updating and accurate layer-wise
    correction. For more information, please refer to `Large Batch Optimization for
    Deep Learning: Training BERT in 76 minutes <https://arxiv.org/abs/1904.00962>`_ .

    The updating of parameters follows:

    ..  math::

        m_t &= \beta_1 m_{t - 1}+ (1 - \beta_1)g_t

        v_t &= \beta_2 v_{t - 1}  + (1 - \beta_2)g_t^2

        m_t &= \frac{m_t}{\beta_1^t}

        v_t &= \frac{v_t}{\beta_2^t}

        r_t &= \frac{m_t}{\sqrt{v_t}+\epsilon}

        w_t &= w_{t-1} -\eta_t \frac{\left \| w_{t-1}\right \|}{\left \| r_t + \lambda w_{t-1}\right \|} (r_t + \lambda w_{t-1})


    where :math:`m` is the 1st moment, and :math:`v` the 2nd moment, :math:`\\eta` the
    learning rate, :math:`\\lambda` the LAMB weight decay rate.

    Args:
        learning_rate (float|Variable, optional): the learning rate used to update parameters. \
            Can be a float value or a Variable with data type float32. Default 0.001.
        lamb_weight_decay (float, optional): The LAMB weight decay rate. Default 0.01. Remind that weight_decay should be None.
        beta1 (float, optional): The exponential decay rate for the 1st moment estimates.
            Default 0.9.
        beta2 (float, optional): The exponential decay rate for the 2nd moment estimates.
            Default 0.999.
        epsilon (float, optional): A small float value for numerical stability. Default 1e-6.
        parameters (Iterable, optional):  Iterable of ``Variable`` names to update to minimize ``loss``. \
            This parameter is required in dygraph mode. And you can specify different options for \
            different parameter groups such as the learning rate, weight decay, etc, \
            then the parameters are list of dict. Note that the learning_rate in parameter groups \
            represents the scale of base learning_rate. \
            The default value is None in static graph mode, at this time all parameters will be updated.
        grad_clip (GradientClipBase, optional): Gradient clipping strategy, it's an instance of
            some derived class of ``GradientClipBase`` . There are three clipping strategies
            ( :ref:`api_paddle_base_clip_ClipGradByGlobalNorm` , :ref:`api_paddle_base_clip_ClipGradByNorm` ,
            :ref:`api_paddle_base_clip_ClipGradByValue` ). If you want better convergence, it is recommended
            to use :ref:`api_paddle_base_clip_ClipGradByGlobalNorm` . Default None, meaning there is no gradient clipping.
        exclude_from_weight_decay_fn (function, optional): whether to skip weight decay for a parameter when this function returns True while take the parameter as input.
        always_adapt (bool, optional): whether to use Layer-wise LR adaptation. By default, skip adaptation on parameters that are
            excluded from weight decay, unless always_adapt == True, then always enable LR adaptation.
        name(str|None): For detailed information, please refer to
            :ref:`api_guide_Name` . Usually name is no need to set and None by default.
    Examples:
        .. code-block:: python

            >>> import paddle

            >>> inp = paddle.uniform(shape=[10, 10], dtype='float32', min=-0.1, max=0.1)
            >>> linear = paddle.nn.Linear(10, 10)
            >>> out = linear(inp)
            >>> loss = paddle.mean(out)
            >>> beta1 = paddle.to_tensor([0.9], dtype="float32")
            >>> beta2 = paddle.to_tensor([0.85], dtype="float32")
            >>> lamb = paddle.optimizer.Lamb(learning_rate=0.002, parameters=linear.parameters(), lamb_weight_decay=0.01)
            >>> back = out.backward()
            >>> lamb.step()
            >>> lamb.clear_grad()

    Úmoment1Úmoment2Úbeta1_pow_accÚbeta2_pow_accc                 óô   •— |€J ‚|€J ‚|€J ‚|€J ‚t         ‰| �  ||d ||¬«       d| _        || _        || _        || _        || _        || _        |||||dœ| _        i | _	        i | _
        |	| _        |
| _        y )N)Úlearning_rateÚ
parametersÚweight_decayÚ	grad_clipÚnameÚlamb)Úbeta1Úbeta2ÚepsilonÚlamb_weight_decayÚexclude_from_weight_decay_fn)ÚsuperÚ__init__ÚtypeÚ_beta1Ú_beta2Ú_epsilonÚ_lamb_weight_decayÚ_exclude_from_weight_decay_fnÚ_default_dictÚ_master_weightsÚ_used_master_weightsÚ_multi_precisionÚalways_adapt)Úselfr   r   r   r   r   r   r   r   Úmulti_precisionr)   r   Ú	__class__s               €ú^G:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/optimizer/lamb.pyr   zLamb.__init__e   s¿   ø€ ð Ð(Ð(Ð(ØÐ Ð Ð ØÐ Ð Ð ØÐ"Ð"Ð"Ü‰ÑØ'Ø!ØØØð 	ô 	
ð ˆŒ	ØˆŒØˆŒØˆŒØ"3ˆÔØ-IˆÔ*àØØØ!2Ø,Hñ
ˆÔð  "ˆÔØ$&ˆÔ!à /ˆÔØ(ˆÕó    c                 óp  — |€
t        «       }|j                  |«      j                  «       }| j                  j	                  |«      }|�i|j                  |«      j                  «       }|j                  «       |j                  «       k7  sJ ‚|j                  «       |j                  «       k(  sJ ‚||fS d }||fS ©N)r   Úfind_varÚ
get_tensorr'   ÚgetÚ_dtypeÚshape)r*   r   ÚscopeÚp_tÚmaster_nameÚ
master_p_ts         r-   Ú_get_parameterzLamb._get_parameter‘   s¯   € Øˆ=Ü “NˆEà�n‰n˜TÓ"×-Ñ-Ó/ˆà×/Ñ/×3Ñ3°DÓ9ˆØÐ"ØŸ™¨Ó4×?Ñ?ÓAˆJØ×$Ñ$Ó&¨#¯*©*«,Ò6Ð6Ð6Ø×#Ñ#Ó%¨¯©«Ò4Ð4Ð4ð �JˆÐð ˆJØ�JˆÐr.   c                 ó  — t        |t        j                  «      sJ ‚t        |t        «      r| j	                  |«      }|D ]À  }|j
                  | j                  v rŒ| j                  rc| j                  |j                  «      rH| j                  |«      }| j                  |«       | j                  j                  |j
                  «       Œ‹| j                  |«       | j                  j                  |j
                  «       ŒÂ y r0   )Ú
isinstancer   ÚBlockÚdictÚ_update_param_groupr   Ú_already_create_accumulaterr(   Ú_is_dtype_fp16_or_bf16ÚdtypeÚ_create_master_weightÚ_add_moments_powsÚadd)r*   Úblockr   ÚpÚmaster_ps        r-   Ú_create_accumulatorszLamb._create_accumulators    sÇ   € Ü˜%¤§¡Ô1Ð1Ð1Ü�j¤$Ô'Ø×1Ñ1°*Ó=ˆJó ˆAØ�v‰v˜×9Ñ9Ñ9ØØ×$Ò$¨×)DÑ)DÀQÇWÁWÔ)MØ×5Ñ5°aÓ8�Ø×&Ñ& xÔ0Ø×0Ñ0×4Ñ4°Q·V±VÕ<à×&Ñ& qÔ)Ø×0Ñ0×4Ñ4°Q·V±VÕ<ñ r.   c           	      óª  — |j                   }| j                  |«      r$t        j                  j                  j
                  }| j                  | j                  ||¬«       | j                  | j                  ||¬«       | j                  | j                  ||t        | j                  t        «      rdn| j                  dgt        j                  j                  j                  d¬«       | j                  | j                  ||t        | j                  t        «      rdn| j                  dgt        j                  j                  j                  d¬«       y )N)rB   çÍÌÌÌÌÌì?r	   Úcpu)r   ÚparamrB   Ú
fill_valuer5   r   Údeviceç+‡ÙÎ÷ï?)rB   rA   r   ÚVarDescÚVarTypeÚFP32Ú_add_accumulatorÚ_moment1_acc_strÚ_moment2_acc_strÚ_beta1_pow_acc_strr<   r    r   Ú
LOD_TENSORÚ_beta2_pow_acc_strr!   )r*   rG   Ú	acc_dtypes      r-   rD   zLamb._add_moments_pows±   s  € Ø—G‘Gˆ	Ø×&Ñ& yÔ1ÜŸ™×,Ñ,×1Ñ1ˆIà×Ñ˜d×3Ñ3°Q¸iÐÔHØ×Ñ˜d×3Ñ3°Q¸iÐÔHØ×ÑØ×(Ñ(ØØä˜$Ÿ+™+¤xÔ0ñ à—‘Ø�#Ü—‘×%Ñ%×0Ñ0Øð 	ô 
	
ð 	×ÑØ×(Ñ(ØØä˜$Ÿ+™+¤xÔ0ñ à—‘Ø�#Ü—‘×%Ñ%×0Ñ0Øð 	õ 
	
r.   c                 óÖ  — t        |t        j                  «      sJ ‚t        |t        «      r| j	                  |«      }d|j
                  _        | j                  | j                  |d   «      }| j                  | j                  |d   «      }| j                  | j                  |d   «      }| j                  | j                  |d   «      }| j                  �| j                  |d   «      rd}n| j                  }| j                  |«      }| j                  xr | j!                  |d   j"                  «      }	|d   j$                  }
|	r)| j&                  |
   }|j$                  | j(                  |
<   nd }t        j*                  «       rRt-        j.                  |d   |d   ||||||d || j0                  | j2                  | j4                  | j6                  |	«       y |d   |d   |||||dœ}|d   ||||dœ}| j0                  | j2                  | j4                  || j6                  |	dœ}|	r
||d<   ||d	<   | j9                  d
«      }|r||d<   |j;                  | j<                  |||d¬«      }|S )NTr   g        r	   )ÚParamÚGradÚLearningRateÚMoment1ÚMoment2ÚBeta1PowÚBeta2Pow)ÚParamOutÚ
Moment1OutÚ
Moment2OutÚBeta1PowOutÚBeta2PowOut)r   r   r   r   r)   r+   ÚMasterParamÚMasterParamOutÚ	found_infÚ
SkipUpdate)r   ÚinputsÚoutputsÚattrsÚstop_gradient)r<   r   r=   r>   r?   ÚprogramÚ	_use_lambÚ_get_accumulator_masterrU   rV   rW   rY   r$   r#   Ú_create_param_lrr(   rA   rB   r   r&   r'   Úin_dygraph_moder   Úlamb_r    r!   r"   r)   Ú_get_auxiliary_varÚ	append_opr   )r*   rF   Úparam_and_gradr   r   r   r   r   ÚlrÚfind_masterÚp_nameÚmaster_weightrl   rm   rn   rj   Úlamb_ops                    r-   Ú_append_optimize_opzLamb._append_optimize_opÏ   sª  € Ü˜%¤§¡Ô1Ð1Ð1Ü�n¤dÔ+Ø!×5Ñ5°nÓEˆNà"&ˆ�‰Ôà×.Ñ.Ø×!Ñ! >°!Ñ#4ó
ˆð ×.Ñ.Ø×!Ñ! >°!Ñ#4ó
ˆð ×4Ñ4Ø×#Ñ# ^°AÑ%6ó
ˆð ×4Ñ4Ø×#Ñ# ^°AÑ%6ó
ˆð
 ×.Ñ.Ð:Ø×2Ñ2°>À!Ñ3DÔEà‰Là×2Ñ2ˆLØ×"Ñ" >Ó2ˆà×+Ñ+ò 
°×0KÑ0KØ˜1Ñ×#Ñ#ó1
ˆð   Ñ"×'Ñ'ˆÙØ ×0Ñ0°Ñ8ˆMØ0=×0BÑ0BˆD×%Ñ% fÒ-à ˆMä×$Ñ$Ô&Ü�L‰LØ˜qÑ!Ø˜qÑ!ØØØØØØØØØ—‘Ø—‘Ø—‘Ø×!Ñ!Øôð" ð (¨Ñ*Ø& qÑ)Ø "Ø"Ø"Ø)Ø)ñˆFð +¨1Ñ-Ø%Ø%Ø,Ø,ñˆGð Ÿ™ØŸ™ØŸ=™=Ø ,Ø $× 1Ñ 1Ø#.ñˆEñ Ø(5��}Ñ%Ø,9�Ð(Ñ)à×/Ñ/°Ó<ˆIÙØ'0��|Ñ$à—o‘oØ—Y‘YØØØØ"ð &ó ˆGð ˆNr.   c                 ó�  — |j                  d| j                  d   «      | _        |j                  d| j                  d   «      | _        |j                  d| j                  d   «      | _        |j                  d| j                  d   «      | _        |j                  d| j                  d   «      | _        |j                  d«      }|S )Nr   r   r   r   r   Úparams)r3   r%   r    r!   r"   r#   r$   )r*   r   s     r-   r?   zLamb._update_param_group6  s¶   € Ø —n‘n W¨d×.@Ñ.@ÀÑ.IÓJˆŒØ —n‘n W¨d×.@Ñ.@ÀÑ.IÓJˆŒØ"Ÿ™ y°$×2DÑ2DÀYÑ2OÓPˆŒØ",§.¡.Ø ×!3Ñ!3Ð4GÑ!Hó#
ˆÔð .8¯^©^Ø*Ø×ÑÐ=Ñ>ó.
ˆÔ*ð  —^‘^ HÓ-ˆ
ØÐr.   )gü©ñÒMbP?g{®Gáz„?rK   rP   g�íµ ÷Æ°>NNNFFNr0   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__rU   rV   rW   rY   r   r:   rI   rD   r~   r?   Ú__classcell__)r,   s   @r-   r   r      sh   ø„ ñEðL !ÐØ ÐØ(ÐØ(Ðð ØØØØØØØ%)ØØØõ*)óXò=ò"
ò<eöNr.   r   N)Úpaddler   Úpaddle.base.executorr   Úbaser   r   Úbase.frameworkr   Ú	optimizerr
   Ú__all__r   © r.   r-   Ú<module>r�      s)   ðõ Ý -ç "Ý %Ý  à
€ôiˆ9õ ir.   