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 )zRAdam Optimizer.
Implementation lifted from: https://github.com/LiyuanLucasLiu/RAdam
Paper: `On the Variance of the Adaptive Learning Rate and Beyond` - https://arxiv.org/abs/1908.03265
    N)	Optimizerc                       s<   e Zd Zd fdd	Z fddZe dd
dZ  ZS )RAdamMbP?g?g+?:0yE>r   c                    s6   t ||||dd tdD d}tt| || d S )Nc                 S   s   g | ]}g d qS ))NNN ).0_r   r   Y/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/radam.py
<listcomp>       z"RAdam.__init__.<locals>.<listcomp>
   )lrbetasepsweight_decaybuffer)dictrangesuperr   __init__)selfparamsr   r   r   r   defaults	__class__r   r
   r      s
    zRAdam.__init__c                    s   t t| | d S )N)r   r   __setstate__)r   stater   r   r
   r      s    zRAdam.__setstate__Nc                 C   s  d }|d ur:t   | }W d    n1 s00    Y  | jD ]x}|d D ]h}|jd u r`qN|j }|jrxtd| }| j| }t|dkrd|d< t 	||d< t 	||d< n$|d 
||d< |d 
||d< |d |d  }}	|d \}
}|	|j||d| d	 ||
j|d|
 d
 |d  d7  < |d t|d d  }|d |d kr~|d |d  }}n|d |d< ||d  }dd|  d }|d|d  | d|   }||d< |dkr$|d td| |d  |d  |d  | | |d   d|
|d    }n|d d|
|d    }||d< |d dkrn|j||d  |d  d
 |dkr|	 |d }|j||| d	 n|j|| d
 || qNq@|S )Nr   z'RAdam does not support sparse gradientsr   stepexp_avg
exp_avg_sqr      )value)alphar   r         r      r   r   )torchZenable_gradZparam_groupsgradfloatZ	is_sparseRuntimeErrorr   lenZ
zeros_likeZtype_asZmul_Zaddcmul_Zadd_intmathsqrtZaddcdiv_Zcopy_)r   closureZlossgrouppr(   Zp_fp32r   r   r    Zbeta1Zbeta2ZbufferedZnum_smaZ	step_sizeZbeta2_tZnum_sma_maxZdenomr   r   r
   r      sx    
$





z
RAdam.step)r   r   r   r   )N)	__name__
__module____qualname__r   r   r'   Zno_gradr   __classcell__r   r   r   r
   r   
   s   r   )__doc__r-   r'   Ztorch.optim.optimizerr   r   r   r   r   r
   <module>   s   