a
    þdc&  ã                   @   s0   d dl Z d dlZd dlmZ G dd„ deƒZdS )é    N)Ú	Optimizerc                	       sP   e Zd ZdZd‡ fdd	„	Z‡ fd
d„Ze ¡ dd„ ƒZe ¡ ddd„ƒZ	‡  Z
S )Ú	AdaBeliefaŠ  Implements AdaBelief algorithm. Modified from Adam in PyTorch

    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-16)
        weight_decay (float, optional): weight decay (L2 penalty) (default: 0)
        amsgrad (boolean, optional): whether to use the AMSGrad variant of this
            algorithm from the paper `On the Convergence of Adam and Beyond`_
            (default: False)
        decoupled_decay (boolean, optional): (default: True) If set as True, then
            the optimizer uses decoupled weight decay as in AdamW
        fixed_decay (boolean, optional): (default: False) This is used when weight_decouple
            is set as True.
            When fixed_decay == True, the weight decay is performed as
            $W_{new} = W_{old} - W_{old} \times decay$.
            When fixed_decay == False, the weight decay is performed as
            $W_{new} = W_{old} - W_{old} \times decay \times lr$. Note that in this case, the
            weight decay ratio decreases with learning rate (lr).
        rectify (boolean, optional): (default: True) If set as True, then perform the rectified
            update similar to RAdam
        degenerated_to_sgd (boolean, optional) (default:True) If set as True, then perform SGD update
            when variance of gradient is high
    reference: AdaBelief Optimizer, adapting stepsizes by the belief in observed gradients, NeurIPS 2020

    For a complete table of recommended hyperparameters, see https://github.com/juntang-zhuang/Adabelief-Optimizer'
    For example train/args for EfficientNet see these gists
      - link to train_scipt: https://gist.github.com/juntang-zhuang/0a501dd51c02278d952cf159bc233037
      - link to args.yaml: https://gist.github.com/juntang-zhuang/517ce3c27022b908bb93f78e4f786dc3
    çü©ñÒMbP?©gÍÌÌÌÌÌì?g+‡ÙÎ÷ï?ç¼‰Ø—²Òœ<r   FTc                    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 s„n t d |d ¡ƒ‚t|ttfƒrüt|ƒdkrüt|d tƒrü|D ]J}d	|v r°|d	 d |d ksä|d	 d |d kr°d
d„ tdƒD ƒ|d< q°t||||||
||	|dd„ tdƒD ƒd
}tt	| ƒ 
||¡ d S )Ng        zInvalid learning rate: {}zInvalid epsilon value: {}r   ç      ð?z%Invalid beta parameter at index 0: {}é   z%Invalid beta parameter at index 1: {}Úbetasc                 S   s   g | ]}g d ¢‘qS ©)NNN© ©Ú.0Ú_r   r   ú]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/adabelief.pyÚ
<listcomp>:   ó    z&AdaBelief.__init__.<locals>.<listcomp>é
   Úbufferc                 S   s   g | ]}g d ¢‘qS r
   r   r   r   r   r   r   ?   r   )
Úlrr	   ÚepsÚweight_decayÚamsgradÚdegenerated_to_sgdÚdecoupled_decayÚrectifyÚfixed_decayr   )Ú
ValueErrorÚformatÚ
isinstanceÚlistÚtupleÚlenÚdictÚrangeÚsuperr   Ú__init__)ÚselfÚparamsr   r	   r   r   r   r   r   r   r   ÚparamÚdefaults©Ú	__class__r   r   r%   *   s$    (0
ýzAdaBelief.__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,   B   s    
zAdaBelief.__setstate__c                 C   sf   | j D ]Z}|d D ]L}| j| }|d }d|d< t |¡|d< t |¡|d< |rt |¡|d< qqd S )Nr'   r   r   ÚstepÚexp_avgÚexp_avg_varÚmax_exp_avg_var)r-   r/   ÚtorchÚ
zeros_like)r&   r0   Úpr/   r   r   r   r   ÚresetG   s    

zAdaBelief.resetNc                 C   s  d}|dur:t  ¡  |ƒ }W d  ƒ n1 s00    Y  | jD ]Æ}|d D ]¶}|jdu r`qN|j}|jt jt jhv r€| ¡ }|jrŽt	dƒ‚|}|jt jt jhv r¬| ¡ }|d }|d \}}	| j
| }
t|
ƒdkrd|
d< t  |¡|
d< t  |¡|
d	< |rt  |¡|
d
< |d rT|d s@| d|d |d   ¡ n| d|d  ¡ n |d dkrt|j||d d |
d |
d	  }}|
d  d7  < d||
d   }d|	|
d   }| |¡j|d| d || }| |	¡j||d|	 d |r6|
d
 }t j|| |d ¡|d | ¡ t |¡  |d ¡}n&| |d ¡ ¡ t |¡  |d ¡}|d sˆ|d | }|j||| d nb|d t|
d d ƒ }|
d |d krÆ|d |d  }}nÊ|
d |d< |	|
d  }dd|	  d }|d|
d  | d|   }||d< |dkrdt d| |d  |d  |d  | | |d  ¡d||
d    }n$|d r„dd||
d    }nd}||d< |dkrÈ| ¡  |d ¡}|j||| |d  d n"|dkrê|j|| |d  d |jt jt jhv rN| |¡ qNq@|S )zµPerforms a single optimization step.
        Arguments:
            closure (callable, optional): A closure that reevaluates the model
                and returns the loss.
        Nr'   zOAdaBelief does not support sparse gradients, please consider SparseAdam insteadr   r	   r   r1   r2   r3   r4   r   r   r   r   r   )Úalphar   )Úvaluer   )Úoutr   r   r   é   é   é   r   éÿÿÿÿ)r5   Zenable_gradr-   ÚgradZdtypeÚfloat16Zbfloat16ÚfloatZ	is_sparseÚRuntimeErrorr/   r!   r6   Zmul_Zadd_Zaddcmul_ÚmaxÚsqrtÚmathZaddcdiv_ÚintZcopy_)r&   ÚclosureZlossr0   r7   r@   Zp_fp32r   Zbeta1Zbeta2r/   r2   r3   Zbias_correction1Zbias_correction2Zgrad_residualr4   ZdenomZ	step_sizeZbufferedZnum_smaZbeta2_tZnum_sma_maxr   r   r   r1   Y   s¬    
$
ÿ


&

ÿÿþþýýÿü


zAdaBelief.step)	r   r   r   r   FTFTT)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r%   r,   r5   Zno_gradr8   r1   Ú__classcell__r   r   r*   r   r      s   $  þ
r   )rF   r5   Ztorch.optim.optimizerr   r   r   r   r   r   Ú<module>   s   