a
    d                     @   s4   d Z ddlZddlZddlmZ G dd deZdS )z AdamW Optimizer
Impl copied from PyTorch master

NOTE: Builtin optim.AdamW is used by the factory, this impl only serves as a Python based reference, will be removed
someday
    N)	Optimizerc                       s@   e Zd ZdZd fdd	Z fd	d
Ze dddZ  Z	S )AdamWa  Implements AdamW algorithm.

    The original Adam algorithm was proposed in `Adam: A Method for Stochastic Optimization`_.
    The AdamW variant was proposed in `Decoupled Weight Decay Regularization`_.

    Arguments:
        params (iterable): iterable of parameters to optimize or dicts defining
            parameter groups
        lr (float, optional): learning rate (default: 1e-3)
        betas (Tuple[float, float], optional): coefficients used for computing
            running averages of gradient and its square (default: (0.9, 0.999))
        eps (float, optional): term added to the denominator to improve
            numerical stability (default: 1e-8)
        weight_decay (float, optional): weight decay coefficient (default: 1e-2)
        amsgrad (boolean, optional): whether to use the AMSGrad variant of this
            algorithm from the paper `On the Convergence of Adam and Beyond`_
            (default: False)

    .. _Adam\: A Method for Stochastic Optimization:
        https://arxiv.org/abs/1412.6980
    .. _Decoupled Weight Decay Regularization:
        https://arxiv.org/abs/1711.05101
    .. _On the Convergence of Adam and Beyond:
        https://openreview.net/forum?id=ryQu7f-RZ
    MbP?g?g+?:0yE>{Gz?Fc                    s   d|kst d|d|ks,t d|d|d   krDdk sXn t d|d d|d   krpdk sn t d|d t|||||d	}tt| || d S )
Ng        zInvalid learning rate: {}zInvalid epsilon value: {}r   g      ?z%Invalid beta parameter at index 0: {}   z%Invalid beta parameter at index 1: {})lrbetasepsweight_decayamsgrad)
ValueErrorformatdictsuperr   __init__)selfparamsr	   r
   r   r   r   defaults	__class__ Y/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/adamw.pyr   '   s    zAdamW.__init__c                    s,   t t| | | jD ]}|dd qd S )Nr   F)r   r   __setstate__param_groups
setdefault)r   stategroupr   r   r   r   5   s    
zAdamW.__setstate__Nc                 C   s  d}|dur:t   | }W d   n1 s00    Y  | jD ]}|d D ]}|jdu r`qN|jd|d |d    |j}|jrtd|d }| j| }t	|dkrd|d	< t 
||d
< t 
||d< |rt 
||d< |d
 |d  }}	|r|d }
|d \}}|d	  d7  < d||d	   }d||d	   }||j|d| d |	|j||d| d |rt j|
|	|
d |
 t| |d }n|	 t| |d }|d | }|j||| d qNq@|S )zPerforms a single optimization step.

        Arguments:
            closure (callable, optional): A closure that reevaluates the model
                and returns the loss.
        Nr   r   r	   r   zJAdam does not support sparse gradients, please consider SparseAdam insteadr   r   stepexp_avg
exp_avg_sqmax_exp_avg_sqr
   )alpha)value)outr   )torchZenable_gradr   graddataZmul_Z	is_sparseRuntimeErrorr   lenZ
zeros_likeZadd_Zaddcmul_maxsqrtmathZaddcdiv_)r   closureZlossr   pr'   r   r   r    r!   r"   Zbeta1Zbeta2Zbias_correction1Zbias_correction2ZdenomZ	step_sizer   r   r   r   :   sH    
$

z
AdamW.step)r   r   r   r   F)N)
__name__
__module____qualname____doc__r   r   r&   Zno_gradr   __classcell__r   r   r   r   r      s     r   )r3   r-   r&   Ztorch.optim.optimizerr   r   r   r   r   r   <module>   s   