Ë
    •\;jmj  ã                   óÎ   — d dl Z d dlmZ d dlmZ d dlZd dlmZ d dlmZ 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mZmZmZ ddlmZ ddlmZ ddlmZ g Z G d„ de«      Zy)é    N)Údefaultdict)ÚCallable)Úpir)ÚDataType)ÚOpResulté   )Ú_C_ops)ÚcoreÚ	framework)Úbase)Ú	ParameterÚVariableÚin_dynamic_or_pir_modeÚin_pir_mode)ÚGradientClipBaseé   )ÚLRScheduler)Ú	Optimizerc                   ó²   — e Zd ZdZdZdZdZdZ	 	 	 	 	 	 	 	 	 	 	 	 dd„Zd„ Z	d	„ Z
d
„ Zd„ Zd„ Zd„ Zd„ Zej"                  ej&                  d„ «       «       Zd„ Zy)ÚAdamWa\  
    The AdamW optimizer is implemented based on the AdamW Optimization
    in paper `DECOUPLED WEIGHT DECAY REGULARIZATION <https://arxiv.org/pdf/1711.05101.pdf>`_.
    it can resolves the problem of L2 regularization failure in the Adam optimizer.

    .. math::

        t & = t + 1

        moment\_1\_out & = {\beta}_1 * moment\_1 + (1 - {\beta}_1) * grad

        moment\_2\_out & = {\beta}_2 * moment\_2 + (1 - {\beta}_2) * grad * grad

        learning\_rate & = learning\_rate *
            \frac{\sqrt{1 - {\beta}_2^t}}{1 - {beta}_1^t}

        param\_out & = param - learning\_rate * (\frac{moment\_1}{\sqrt{moment\_2} + \epsilon} + \lambda * param)


    Args:
        learning_rate (float|LRScheduler, optional): The learning rate used to update ``Parameter``.
            It can be a float value or a LRScheduler. The default value is 0.001.
        parameters (list|tuple, optional): List/Tuple of ``Tensor`` 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.
        beta1 (float|Tensor, optional): The exponential decay rate for the 1st moment estimates.
            It should be a float number or a Tensor with shape [1] and data type as float32.
            The default value is 0.9.
        beta2 (float|Tensor, optional): The exponential decay rate for the 2nd moment estimates.
            It should be a float number or a Tensor with shape [1] and data type as float32.
            The default value is 0.999.
        epsilon (float, optional): A small float value for numerical stability.
            The default value is 1e-08.
        weight_decay (float|Tensor, optional): The weight decay coefficient, it can be float or Tensor. The default value is 0.01.
        lr_ratio (function|None, optional): If it is not None,
            the learning rate will be updated with layer-wise learning rate ratio.
            Otherwise, the learning rate is the original.
            Default: None.
        apply_decay_param_fun (function|None, optional): If it is not None,
            only tensors that makes apply_decay_param_fun(Tensor.name)==True
            will be updated with weight decay. It only works when we want to specify tensors.
            Default: None.
        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_nn_ClipGradByGlobalNorm` , :ref:`api_paddle_nn_ClipGradByNorm` ,
            :ref:`api_paddle_nn_ClipGradByValue` ). Default None, meaning there is no gradient clipping.
        lazy_mode (bool, optional): The official Adam algorithm has two moving-average accumulators.
            The accumulators are updated at every step. Every element of the two moving-average
            is updated in both dense mode and sparse mode. If the size of parameter is very large,
            then the update may be very slow. The lazy mode only update the element that has
            gradient in current mini-batch, so it will be much more faster. But this mode has
            different semantics with the original Adam algorithm and may lead to different result.
            The default value is False.
        multi_precision (bool, optional): Whether to use multi-precision during weight updating. Default is false.
        name (str, optional): Normally there is no need for user to set this property.
            For more information, please refer to :ref:`api_guide_Name`.
            The default value is None.
    Notes:
        **Currently, AdamW doesn't support sparse parameter optimization.**

    Examples:
        .. code-block:: python

            >>> import paddle

            >>> linear = paddle.nn.Linear(10, 10)
            >>> inp = paddle.rand([10,10], dtype="float32")
            >>> out = linear(inp)
            >>> loss = paddle.mean(out)

            >>> beta1 = paddle.to_tensor([0.9], dtype="float32")
            >>> beta2 = paddle.to_tensor([0.99], dtype="float32")

            >>> opt = paddle.optimizer.AdamW(learning_rate=0.1,
            ...         parameters=linear.parameters(),
            ...         beta1=beta1,
            ...         beta2=beta2,
            ...         weight_decay=0.01
            ... )
            >>> loss.backward()
            >>> opt.step()
            >>> opt.clear_grad()


            >>> # Note that the learning_rate of linear_2 is 0.01.
            >>> linear_1 = paddle.nn.Linear(10, 10)
            >>> linear_2 = paddle.nn.Linear(10, 10)
            >>> inp = paddle.uniform(shape=[10, 10], min=-0.1, max=0.1)
            >>> out = linear_1(inp)
            >>> out = linear_2(out)
            >>> loss = paddle.mean(out)
            >>> opt = paddle.optimizer.AdamW(
            ...     learning_rate=0.1,
            ...     parameters=[{
            ...         'params': linear_1.parameters()
            ...     }, {
            ...         'params': linear_2.parameters(),
            ...         'weight_decay': 0.001,
            ...         'learning_rate': 0.1,
            ...         'beta1': 0.8
            ...     }],
            ...     weight_decay=0.01,
            ...     beta1=0.9
            ... )
            >>> loss.backward()
            >>> opt.step()
            >>> opt.clear_grad()

    Úmoment1Úmoment2Úbeta1_pow_accÚbeta2_pow_accNc                 óR  — |€J ‚|€J ‚|€J ‚|€J ‚d|cxk  rdk  st        d«      ‚ t        d«      ‚d|cxk  rdk  st        d«      ‚ t        d«      ‚d|k  st        d«      ‚t        |t        «      s+t        |t        j                  t
        f«      st        d«      ‚|�“t        |t        «      sJ ‚t        j                  «       smt        j                  «       sYt        j                  j                  «       j                  d«      d   t        j                  j                  «       vrt!        d«      ‚|�ƒt        |t        j"                  t        j$                  j"                  f«      r#t        d	j'                  t)        |«      «      «      ‚t        |t*        «      rt        d
«      ‚t-        |«      | _        nd | _        || _        t        j2                  «       r| j.                  €t5        d«      ‚t        |t        t6        f«      st        dt)        |«      z  «      ‚|	�t        |	t8        «      st        d«      ‚d | _        | j.                  r|t        | j.                  d   t*        «      rA| j.                  D ]  }d|v rŒJ d«       ‚ | j.                  d   d   d   j<                  | _        n| j.                  d   j<                  | _        i | _        tA        d„ «      | _!        d | _"        g | _#        i | _$        i | _%        | jL                  | _'        d| _        || _(        tS        «       | _*        || _+        || _,        |	| _-        || _.        || _/        || _0        || _1        |
| _2        || _3        i | _4        |||||
|	dœ| _5        g | _6        | j.                  rNt        | j.                  d   t*        «      r1| j.                  D ]!  }| jo                  |jq                  «       «       Œ# n| j.                  | _6        d | _9        d | _:        i | _;        tS        «       | _<        | j{                  «        y )Nr   r   z.Invaild value of beta1, expect beta1 in [0,1).z.Invaild value of beta2, expect beta2 in [0,1).z.Invaild value of epsilon, expect epsilon >= 0.z'weight_decay should be float or Tensor.Ú:z#'lr_ratio' is unimplemented in CPU.zt`parameters` argument given to the optimizer should be an iterable of paddle Tensors, but got argument type is `{}`.zv`parameters` argument should not get dict type, if parameter groups is needed, please set `parameters` as list of dictzNparameters argument given to the Optimizer should not be None in dygraph mode.z9learning rate should be float or LRScheduler, got %s herezE'grad_clip' should be an instance of GradientClipBase's derived classÚparamszYparams should be set in parameters if parameter groups are optimized in different optionsc                  ó   — i S ©N© r    ó    ú_G:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/optimizer/adamw.pyÚ<lambda>z AdamW.__init__.<locals>.<lambda>   s   € ±r!   Úadamw)Úweight_decayÚbeta1Úbeta2ÚepsilonÚ	lazy_modeÚ	grad_clip)>Ú
ValueErrorÚ
isinstanceÚfloatr   r   r   Ú	TypeErrorr   r
   Úis_compiled_with_cudaÚis_compiled_with_xpuÚpaddleÚdeviceÚ
get_deviceÚsplitÚget_all_custom_device_typeÚNotImplementedErrorÚTensorÚeagerÚformatÚtypeÚdictÚlistÚ_parameter_listÚ_nameÚin_dygraph_modeÚAttributeErrorr   r   Ú_dtypeÚdtypeÚ_learning_rate_mapr   Ú_accumulatorsÚhelperÚ_opti_name_listÚ_accumulators_holderÚ_param_device_mapÚ
clear_gradÚclear_gradientsÚ_learning_rateÚsetÚ_params_nameÚ_apply_decay_param_funÚ_weight_decayÚ
_grad_clipÚ	_lr_ratioÚ_beta1Ú_beta2Ú_epsilonÚ
_lazy_modeÚ_multi_precisionÚ_master_weightsÚ_default_dictÚ_param_groupsÚ_add_param_groupÚcopyÚ_use_multi_tensorÚregularizationÚ_auxiliary_varsÚ_already_create_accumulaterÚ_create_master_grad_states)ÚselfÚlearning_rater&   r'   r(   Ú
parametersr%   Úlr_ratioÚapply_decay_param_funr*   r)   Úmulti_precisionÚnameÚparam_groups                 r"   Ú__init__zAdamW.__init__Ÿ   s  € ð Ð(Ð(Ð(ØÐ Ð Ð ØÐ Ð Ð ØÐ"Ð"Ð"Ø�EŒ~˜AŠ~ÜÐMÓNÐNð ÜÐMÓNÐNØ�EŒ~˜AŠ~ÜÐMÓNÐNð ÜÐMÓNÐNØ�GŠ|ÜÐMÓNÐNÜ˜,¬Ô.´zØœ9×-Ñ-¬xÐ8ô8
ô ÐEÓFÐFØÐÜ˜h¬Ô1Ð1Ð1ä×.Ñ.Ô0Ü×1Ñ1Ô3Ü—M‘M×,Ñ,Ó.×4Ñ4°SÓ9¸!Ñ<Ü—}‘}×?Ñ?ÓAñBô *Ð*OÓPÐPàÐ!ô ˜*¤v§}¡}´d·j±j×6GÑ6GÐ&HÔIÜðTßTZÑTZÜ˜ZÓ(óUóð ô ˜*¤dÔ+Üð'óð ô
 $(¨
Ó#3ˆDÕ à#'ˆDÔ àˆŒ
Ü×$Ñ$Ô&Ø×#Ñ#Ð+Ü$Ødóð ô ˜-¬%´Ð)=Ô>ÜØKÜ�}Ó%ñ&óð ð Ð Ü˜iÔ)9Ô:ÜØ[óð ð ˆŒà×ÒÜ˜$×.Ñ.¨qÑ1´4Ô8Ø#'×#7Ô#7�Kà  KÒ/ðsàrósØ/ð $8ð #×2Ñ2°1Ñ5°hÑ?ÀÑB×HÑH�•à"×2Ñ2°1Ñ5×;Ñ;�”ð #%ˆÔô
 )©Ó4ˆÔØˆŒØ!ˆÔØ$&ˆÔ!Ø!#ˆÔØ#Ÿ™ˆÔàˆŒ	Ø+ˆÔÜ›EˆÔØ&;ˆÔ#Ø)ˆÔØ#ˆŒØ!ˆŒØˆŒØˆŒØˆŒØ#ˆŒØ /ˆÔØ!ˆÔð )ØØØØ"Ø"ñ
ˆÔð  ˆÔØ×Ò¤J¨t×/CÑ/CÀAÑ/FÌÔ$MØ#×3Ô3�Ø×%Ñ% k×&6Ñ&6Ó&8Õ9ñ  4ð "&×!5Ñ!5ˆDÔà!%ˆÔØ"ˆÔØ!ˆÔÜ+.«5ˆÔ(à×'Ñ'Õ)r!   c                 ó"   — || j                   |<   y r   ©r^   )ra   ÚkeyÚvals      r"   Ú_set_auxiliary_varzAdamW._set_auxiliary_var,  s   € Ø$'ˆ×Ñ˜SÒ!r!   c                 ó>   — || j                   v r| j                   |   S y r   rk   )ra   rl   s     r"   Ú_get_auxiliary_varzAdamW._get_auxiliary_var/  s$   € Ø�$×&Ñ&Ñ&Ø×'Ñ'¨Ñ,Ð,àr!   c                 ór  — |d   }t        |t        t        j                  j                  f«      r|g|d<   n)t        |t
        «      rt        d«      ‚t        |«      |d<   | j                  j                  «       D ]  \  }}|j                  ||«       Œ t        «       }| j                  D ]  }|j                  t        |d   «      «       Œ! |j                  t        |d   «      «      st        d«      ‚|d   D ]!  }|j                  dd«      |j                   d<   Œ# | j                  j#                  |«       y)zº
        Add a param group to parameter_list.

        Args:
            param_group (dict): The group of Tensors to be optimzed with
            different optimization options.
        r   z`optimizer parameters should be in ordered collections,but received set, please use list instead.z7some parameters appear in more than one parameter grouprb   ç      ð?N)r,   r   r   r
   ÚParameterMetarL   r.   r<   rX   ÚitemsÚ
setdefaultrY   ÚupdateÚ
isdisjointr+   ÚgetÚoptimize_attrÚappend)ra   rh   r   ÚkÚvÚ	param_setÚgroupÚparams           r"   rZ   zAdamW._add_param_group5  s-  € ð ˜XÑ&ˆÜ�fœy¬#¯(©(×*@Ñ*@ÐAÔBØ%+ HˆK˜Ò!Ü˜¤Ô$Üð=óð ô
 %)¨£LˆK˜Ñ!ð ×&Ñ&×,Ñ,Ö.‰DˆAˆqØ×"Ñ" 1 aÕ(ð /ô “Eˆ	Ø×'Ô'ˆEØ×ÑœS  x¡Ó1Õ2ð (ð ×#Ñ#¤C¨°HÑ(=Ó$>Ô?ÜØIóð ð ! Ô*ˆEØ3>·?±?Ø ó4ˆE×Ñ Ò0ð +ð
 	×Ñ×!Ñ! +Õ.r!   c           
      óÒ  — |j                   }| j                  |«      r>t        «       rt        j                  n#t
        j                  j                  j                  }t        j                  «       rÚdd l
}|j                  dd¬«      }|dk(  r�| j                  | j                  |t
        j                  j                  j                  ¬«       | j                  | j                  |t
        j                  j                  j                  ¬«       ny| j                  | j                  ||¬«       | j                  | j                  ||¬«       n<| j                  | j                  ||¬«       | j                  | j                  ||¬«       | j                  | j                   ||t#        | j$                  t&        t(        f«      rdn| j$                  dgt
        j                  j                  j*                  d	¬
«       | j                  | j,                  ||t#        | j.                  t&        t(        f«      rdn| j.                  dgt
        j                  j                  j*                  d	¬
«       y )Nr   Úxpu_adamw_moment_dtypeÚfp32)ÚdefaultÚfp16)rB   çÍÌÌÌÌÌì?r   Úcpu)rg   r   rB   Ú
fill_valueÚshaper:   r2   ç+‡ÙÎ÷ï?)rB   Ú_is_dtype_fp16_or_bf16r   r   ÚFLOAT32r
   ÚVarDescÚVarTypeÚFP32r0   ÚosÚgetenvÚ_add_accumulatorÚ_moment1_acc_strÚFP16Ú_moment2_acc_strÚ_beta1_pow_acc_strr,   rR   r   r   Ú
LOD_TENSORÚ_beta2_pow_acc_strrS   )ra   ÚpÚ	acc_dtyper�   r�   s        r"   Ú_add_moments_powszAdamW._add_moments_pows\  sû  € Ø—G‘Gˆ	Ø×&Ñ& yÔ1ä$/¤M”× Ò ´t·|±|×7KÑ7K×7PÑ7Pð ô ×$Ñ$Ô&Ûà%'§Y¡YØ(°&ð &/ó &Ð"ð &¨Ò/Ø×%Ñ%Ø×)Ñ)¨1´D·L±L×4HÑ4H×4MÑ4Mð &ô ð ×%Ñ%Ø×)Ñ)¨1´D·L±L×4HÑ4H×4MÑ4Mð &õ ð ×%Ñ% d×&;Ñ&;¸QÀiÐ%ÔPØ×%Ñ% d×&;Ñ&;¸QÀiÐ%ÕPà×!Ñ! $×"7Ñ"7¸À)Ð!ÔLØ×!Ñ! $×"7Ñ"7¸À)Ð!ÔLØ×ÑØ×(Ñ(ØØä˜$Ÿ+™+¬´(Ð';Ô<ñ à—‘Ø�#Ü—‘×%Ñ%×0Ñ0Øð 	ô 
	
ð 	×ÑØ×(Ñ(ØØä˜$Ÿ+™+¬´(Ð';Ô<ñ à—‘Ø�#Ü—‘×%Ñ%×0Ñ0Øð 	õ 
	
r!   c                 ó   — t        |t        j                  t        j                  f«      sJ ‚t        |t        «      r| j                  |«      }|D ]ü  }|j                  | j                  v rŒ| j                  rc| j                  |j                  «      rH| j                  |«      }| j                  |«       | j                  j                  |j                  «       Œ‹| j                  |j                  «      r!| j                  st        j                  d«       | j                  |«       | j                  j                  |j                  «       Œþ y )Nz›Accumulating with FP16 or BF16 in optimizer can lead to poor accuracy or slow convergence.Consider using multi_precision=True option of the Adam optimizer.)r,   r   ÚBlockr   r;   Ú_update_param_grouprg   r_   rV   rŠ   rB   Ú_create_master_weightrš   ÚaddÚwarningsÚwarn)ra   Úblockrc   r˜   Úmaster_ps        r"   Ú_create_accumulatorszAdamW._create_accumulatorsŒ  s  € Ü˜%¤)§/¡/´3·9±9Ð!=Ô>Ð>Ð>Ü�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Ô<Øà×+Ñ+¨A¯G©GÔ4Ø×-Ò-ä—‘ðXôð ×"Ñ" 1Ô%Ø×,Ñ,×0Ñ0°·±Õ8ñ# r!   c                 óZ  — t        |t        j                  t        j                  f«      sJ ‚t        |t        «      r| j                  |«      }|\  }}d}| j                  �| j                  |j                  «      sd}| j                  | j                  |d   «      }| j                  | j                  |d   «      }| j                  | j                  |d   «      }| j                  | j                  |d   «      }	| j                  xr | j                  |d   j                  «      }
|
r| j                   |d   j                     nd }| j#                  |«      }t%        «       r÷| j&                  €dn| j'                  |d   «      }t        | j(                  t*        «      s| j(                  n| j(                  j-                  d«      }t        | j.                  t*        «      s| j.                  n| j.                  j-                  d«      }t1        j2                  |d   |d   |||||	|d ||| j4                  || j6                  || j8                  d|
d«      \  }}}}}}y |d   g|d   g|g|g|g|g|	gdœ}| j;                  d«      }|r||d	<   |d   g|g|g|g|	gd
œ}| j8                  d|
|| j6                  | j&                  €dn| j'                  |d   «      dœ}t        | j(                  t*        «      r| j(                  |d<   n| j(                  |d<   t        | j.                  t*        «      r| j.                  |d<   n| j.                  |d<   t        | j4                  t*        «      r| j4                  |d<   n| j4                  |d<   |
r
||d<   ||d<   |j=                  | j>                  |||d¬«      }|S )NTFr   rr   r   iè  )ÚParamÚGradÚLearningRateÚMoment1ÚMoment2ÚBeta1PowÚBeta2PowÚ	found_infÚ
SkipUpdate)ÚParamOutÚ
Moment1OutÚ
Moment2OutÚBeta1PowOutÚBeta2PowOut)r)   Úmin_row_size_to_use_multithreadrf   Ú
with_decayÚcoeffrd   ÚBeta1Tensorr&   ÚBeta2Tensorr'   ÚEpsilonTensorr(   ÚMasterParamÚMasterParamOut)r:   ÚinputsÚoutputsÚattrsÚstop_gradient) r,   r   rœ   r   r;   r�   rN   rg   Ú_get_accumulator_masterr’   r”   r•   r—   rV   rŠ   rB   rW   Ú_create_param_lrr   rQ   rR   r   ÚitemrS   r	   Úadamw_rT   rO   rU   rp   Ú	append_opr:   )ra   r¢   Úparam_and_gradr   Úgradrµ   r   r   r   r   Úfind_masterÚmaster_weightÚlrÚ	lr_ratio_rR   rS   Ú_r¼   r­   r½   r¾   Úadamw_ops                         r"   Ú_append_optimize_opzAdamW._append_optimize_op¥  sß  € Ü˜%¤)§/¡/´3·9±9Ð!=Ô>Ð>Ð>Ü�n¤dÔ+Ø!×5Ñ5°nÓEˆNØ$‰ˆˆtð ˆ
à×'Ñ'Ð3Ø×/Ñ/°·
±
Ô;àˆJà×.Ñ.Ø×!Ñ! >°!Ñ#4ó
ˆð ×.Ñ.Ø×!Ñ! >°!Ñ#4ó
ˆð ×4Ñ4Ø×#Ñ# ^°AÑ%6ó
ˆð ×4Ñ4Ø×#Ñ# ^°AÑ%6ó
ˆð ×+Ñ+ò 
°×0KÑ0KØ˜1Ñ×#Ñ#ó1
ˆñ
 ð × Ñ  °Ñ!2×!7Ñ!7Ò8àð 	ð
 ×"Ñ" >Ó2ˆô "Ô#ð —>‘>Ð)ñ à—^‘^ N°1Ñ$5Ó6ð ô " $§+¡+¬xÔ8ð —’à—[‘[×%Ñ% aÓ(ð ô " $§+¡+¬xÔ8ð —’à—[‘[×%Ñ% aÓ(ð ô  &Ÿ}™}Ø˜qÑ!Ø˜qÑ!ØØØØØØØØØØ—‘ØØ×"Ñ"ØØ—‘ØØØó' ÑˆAˆq�!�Q˜˜1ð* ð )¨Ñ+Ð,Ø'¨Ñ*Ð+Ø!# Ø#˜9Ø#˜9Ø*˜OØ*˜OñˆFð ×/Ñ/°Ó<ˆIáØ'0��|Ñ$ð ,¨AÑ.Ð/Ø&˜iØ&˜iØ -˜Ø -˜ñˆGð "Ÿ_™_Ø37Ø#.Ø(Ø×+Ñ+à—>‘>Ð)ñ  à—^‘^ N°1Ñ$5Ó6ñ	ˆEô ˜$Ÿ+™+¤xÔ0Ø(,¯©��}Ò%à!%§¡��g‘Ü˜$Ÿ+™+¤xÔ0Ø(,¯©��}Ò%à!%§¡��g‘Ü˜$Ÿ-™-¬Ô2Ø*.¯-©-��Ò'à#'§=¡=��iÑ áØ(5��}Ñ%Ø,9�Ð(Ñ)à—‘Ø—Y‘YØØØØ"ð 'ó ˆHð ˆOr!   c                 óZ   — dj                  ddj                  | j                  «      g«      S )NÚ zWeight Decay, params:Ú,)ÚjoinrM   )ra   s    r"   Ú__str__zAdamW.__str__0  s&   € Ø�x‰xÐ0°#·(±(¸4×;LÑ;LÓ2MÐNÓOÐOr!   c           	      óþ  — t         j                  j                  j                  j                  «       r| j	                  «        yt        | j                  d   t        «      sãg }| j                  D ]½  }|j                  rŒ|j                  «       €Œ!|j                  «       }t        j                  «       r3t        |d«      rZ|j                  «       rJ| j                  �>t        d«      ‚t        |d«      r'|j!                  «       r| j                  �t        d«      ‚|j#                  ||f«       Œ¿ | j%                  dd|¬«      }y| j&                  D �]$  }t)        d„ «      }|d   D ]À  }|j                  rŒ|j                  «       €Œ!|j                  «       }t        j                  «       r3t        |d«      rZ|j                  «       rJ| j                  �>t        d«      ‚t        |d«      r'|j!                  «       r| j                  �t        d«      ‚|d   j#                  ||f«       ŒÂ |j+                  |j-                  «       D ��ci c]  \  }}|dk7  sŒ||“Œ c}}«       | j%                  dd|¬«       �Œ' yc c}}w )	aœ  
        Execute the optimizer and update parameters once.

        Returns:
            None

        Examples:
            .. code-block:: python

                >>> import paddle

                >>> a = paddle.rand([2,13], dtype="float32")
                >>> linear = paddle.nn.Linear(13, 5)
                >>> # This can be any optimizer supported by dygraph.
                >>> opt = paddle.optimizer.AdamW(learning_rate = 0.01,
                ...                             parameters = linear.parameters())
                >>> out = linear(a)
                >>> out.backward()
                >>> opt.step()
                >>> opt.clear_grad()
        Nr   Úis_selected_rowszOAdamW don't support weight_decay with sparse parameters, please set it to None.Ú
_is_sparse)ÚlossÚstartup_programÚparams_gradsc                  ó   — g S r   r    r    r!   r"   r#   zAdamW.step.<locals>.<lambda>p  s   € ±2r!   r   )r1   r   ÚdygraphÚin_to_static_modeÚ_declarative_stepr,   r=   r;   r¿   Ú
_grad_ivarr   r?   ÚhasattrrÔ   r]   ÚRuntimeErrorrÕ   rz   Ú_apply_optimizerY   r   rv   rt   )ra   rØ   r   Úgrad_varÚoptimize_opsrh   r{   r|   s           r"   Ústepz
AdamW.step3  sf  € ô0 �;‰;×Ñ×#Ñ#×5Ñ5Ô7Ø×"Ñ"Ô$Øä˜$×.Ñ.¨qÑ1´4Ô8ØˆLØ×-Ô-�Ø×&Ò&ØØ×#Ñ#Ó%Ñ1Ø$×/Ñ/Ó1�HÜ ×0Ñ0Ô2ä# HÐ.@ÔAØ (× 9Ñ 9Ô ;Ø $× 3Ñ 3Ð ?ä".Ø qó#ð ô
 $ H¨lÔ;Ø (× 3Ñ 3Ô 5Ø $× 3Ñ 3Ð ?ä".Ø qó#ð ð !×'Ñ'¨°Ð(9Õ:ð/ .ð2  ×/Ñ/Ø¨4¸lð 0ó ‰Lð
  $×1Õ1�Ü*©:Ó6�Ø(¨Ô2�EØ×*Ò*Ø Ø×'Ñ'Ó)Ñ5Ø#(×#3Ñ#3Ó#5˜Ü$×4Ñ4Ô6ä '¨Ð2DÔ EØ$,×$=Ñ$=Ô$?Ø$(×$7Ñ$7Ð$Cä&2Ø$uó'"ð !"ô
 !(¨°,Ô ?Ø$,×$7Ñ$7Ô$9Ø$(×$7Ñ$7Ð$Cä&2Ø$uó'"ð !"ð % XÑ.×5Ñ5°u¸hÐ6GÕHð/ 3ð0 ×#Ñ#Ø&1×&7Ñ&7Ô&9ÔKÑ&9™d˜a ¸QÀ(»]�Q˜‘TÐ&9ÒKôð ×$Ñ$Ø¨tÀ,ð %ö ñ;  2ùó6 Ls   ÉI9ÉI9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%   r   )rx   rX   rR   rS   rT   rU   rO   )ra   rc   s     r"   r�   zAdamW._update_param_group�  s²   € Ø —n‘n W¨d×.@Ñ.@ÀÑ.IÓJˆŒØ —n‘n W¨d×.@Ñ.@ÀÑ.IÓJˆŒØ"Ÿ™ y°$×2DÑ2DÀYÑ2OÓPˆŒØ$Ÿ.™.Ø˜×+Ñ+¨KÑ8ó
ˆŒð (Ÿ^™^Ø˜D×.Ñ.¨~Ñ>ó
ˆÔð  —^‘^ HÓ-ˆ
àÐr!   )gü©ñÒMbP?r…   r‰   g:Œ0âŽyE>Ng{®Gáz„?NNNFFN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r’   r”   r•   r—   ri   rn   rp   rZ   rš   r¤   rÍ   rÒ   Úimperative_baseÚno_gradr   Únon_static_onlyrã   r�   r    r!   r"   r   r   (   s¨   „ ñoðb !ÐØ ÐØ(ÐØ(Ðð ØØØØØØØ"ØØØØóK*òZ(òò%/òN.
ò`9ò2IòVPð ×ÑØ×ÑñYó ó ðYóvr!   r   )r    Úcollectionsr   Úcollections.abcr   r1   r   Úpaddle.base.libpaddler   Ú
paddle.pirr   Ú r	   r   r
   r   Úbase.dygraphré   Úbase.frameworkr   r   r   r   Únn.clipr   rÉ   r   Ú	optimizerr   Ú__all__r   r    r!   r"   Ú<module>rö      sM   ðó Ý #Ý $ã Ý Ý *Ý å ß "Ý 2÷ó õ 'Ý Ý  à
€ôt	ˆIõ t	r!   