a
    þd#  ã                   @   s,   d Z ddlZddlZG dd„ dejjƒZdS )z Adafactor Optimizer

Lifted from https://github.com/pytorch/fairseq/blob/master/fairseq/optim/adafactor.py

Original header/copyright below.

é    Nc                	       s`   e Zd ZdZd‡ fd
d„	Zedd„ ƒZedd„ ƒZedd„ ƒZdd„ Z	e
 ¡ ddd„ƒZ‡  ZS )Ú	Adafactora£  Implements Adafactor algorithm.
    This implementation is based on: `Adafactor: Adaptive Learning Rates with Sublinear Memory Cost`
    (see https://arxiv.org/abs/1804.04235)

    Note that this optimizer internally adjusts the learning rate depending on the
    *scale_parameter*, *relative_step* and *warmup_init* options.

    To use a manual (external) learning rate schedule you should set `scale_parameter=False` and
    `relative_step=False`.

    Arguments:
        params (iterable): iterable of parameters to optimize or dicts defining parameter groups
        lr (float, optional): external learning rate (default: None)
        eps (tuple[float, float]): regularization constants for square gradient
            and parameter scale respectively (default: (1e-30, 1e-3))
        clip_threshold (float): threshold of root mean square of final gradient update (default: 1.0)
        decay_rate (float): coefficient used to compute running averages of square gradient (default: -0.8)
        beta1 (float): coefficient used for computing running averages of gradient (default: None)
        weight_decay (float, optional): weight decay (L2 penalty) (default: 0)
        scale_parameter (bool): if True, learning rate is scaled by root mean square of parameter (default: True)
        warmup_init (bool): time-dependent learning rate computation depends on
            whether warm-up initialization is being used (default: False)
    Nç ÂëþKH´9çü©ñÒMbP?ç      ð?çš™™™™™é¿ç        TFc                    s\   | }|
r|st dƒ‚|d u r"d n|d }t||||||||	||
d
}tt| ƒ ||¡ d S )Nz'warmup_init requires relative_step=Truer   )
ÚlrÚepsÚ	eps_scaleÚclip_thresholdÚ
decay_rateÚbeta1Úweight_decayÚscale_parameterÚrelative_stepÚwarmup_init)Ú
ValueErrorÚdictÚsuperr   Ú__init__)ÚselfÚparamsr   r	   r
   r   r   Zbetasr   r   r   r   r   Údefaults©Ú	__class__© ú]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/adafactor.pyr   )   s    þzAdafactor.__init__c                 C   sj   | d rb| d rd|d  nd}t |dt |d ¡ ƒ}d}| d rVt| d |d	 ƒ}|| | d
< | d
 S )Nr   r   gíµ ÷Æ°>Ústepg{®Gáz„?r   r   r
   ÚRMSr   )ÚminÚmathÚsqrtÚmax)Úparam_groupZparam_stateZmin_stepÚlr_tZparam_scaler   r   r   Ú_get_lr5   s    zAdafactor._get_lrc                 C   s    t |ƒdk}| d d u}||fS )Né   r   )Úlen)r#   Zparam_shapeÚfactoredÚuse_first_momentr   r   r   Ú_get_options@   s    zAdafactor._get_optionsc                 C   s   |   d¡|  ¡ d  S )Nr&   g      à?)ZnormZnumel)Ztensorr   r   r   Ú_rmsF   s    zAdafactor._rmsc                 C   s6   ||j ddd  ¡  d¡}| d¡ ¡ }t ||¡S )NéÿÿÿÿT)ÚdimZkeepdiméþÿÿÿ)ÚmeanZrsqrt_Z	unsqueezeÚrsqrtÚtorchÚmul)r   Úexp_avg_sq_rowÚexp_avg_sq_colZr_factorZc_factorr   r   r   Ú_approx_sq_gradJ   s    zAdafactor._approx_sq_gradc                 C   sR  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
| }|  ||j¡\}}t|ƒdkr<d|d< |rÒt  |¡|d< |r$t  |jdd… ¡ |¡|d< t  |jdd	… |jdd…  ¡ |¡|d
< nt  |¡|d< d|d< nV|rT|d  |¡|d< |r€|d  |¡|d< |d
  |¡|d
< n|d  |¡|d< |}	|jt jt jhv r²|	 ¡ }	|d  d7  < |  |	¡|d< |  ||¡}
dt |d |d ¡ }|d |d  }|rr|d }|d
 }| |¡j|jddd| d | |¡j|jd	dd| d |  ||¡}| |¡ n.|d }| |¡j|d| d | ¡  |¡}| |  |¡|d  jdd¡ | |
¡ |rü|d }| |d ¡j|d|d  d |}|d 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   z,Adafactor does not support sparse gradients.r   r   Úexp_avgr,   r3   r.   r4   Ú
exp_avg_sqr   é   r   r   r&   r	   )r-   )Úalphar   )r   r   r   )r1   Zenable_gradZparam_groupsÚgradZdtypeÚfloat16Zbfloat16ÚfloatZ	is_sparseÚRuntimeErrorÚstater*   Úshaper'   Z
zeros_likeÚzerosÚtor+   r%   r    ÚpowZmul_Zadd_r/   r5   r0   Zdiv_Zclamp_Zcopy_)r   ÚclosureZlossÚgroupÚpr:   r>   r(   r)   Zp_fp32r$   Zbeta2tÚupdater3   r4   r7   r6   r   r   r   r   O   sx    
$

.
   
 zAdafactor.step)	Nr   r   r   r   Nr   TF)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ústaticmethodr%   r*   r+   r5   r1   Zno_gradr   Ú__classcell__r   r   r   r   r      s     ÿ



r   )rJ   r1   r    ZoptimZ	Optimizerr   r   r   r   r   Ú<module>   s   