a
    d                     @   s,   d Z ddlZddlmZ G dd deZdS )aL   RMSProp modified to behave like Tensorflow impl

Originally cut & paste from PyTorch RMSProp
https://github.com/pytorch/pytorch/blob/063946d2b3f3f1e953a2a3b54e0b34f1393de295/torch/optim/rmsprop.py
Licensed under BSD-Clause 3 (ish), https://github.com/pytorch/pytorch/blob/master/LICENSE

Modifications Copyright 2021 Ross Wightman
    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 )	RMSpropTFaE  Implements RMSprop algorithm (TensorFlow style epsilon)

    NOTE: This is a direct cut-and-paste of PyTorch RMSprop with eps applied before sqrt
    and a few other modifications to closer match Tensorflow for matching hyper-params.

    Noteworthy changes include:
    1. Epsilon applied inside square-root
    2. square_avg initialized to ones
    3. LR scaling of update accumulated in momentum buffer

    Proposed by G. Hinton in his
    `course <http://www.cs.toronto.edu/~tijmen/csc321/slides/lecture_slides_lec6.pdf>`_.

    The centered version first appears in `Generating Sequences
    With Recurrent Neural Networks <https://arxiv.org/pdf/1308.0850v5.pdf>`_.

    Arguments:
        params (iterable): iterable of parameters to optimize or dicts defining
            parameter groups
        lr (float, optional): learning rate (default: 1e-2)
        momentum (float, optional): momentum factor (default: 0)
        alpha (float, optional): smoothing (decay) constant (default: 0.9)
        eps (float, optional): term added to the denominator to improve
            numerical stability (default: 1e-10)
        centered (bool, optional) : if ``True``, compute the centered RMSProp,
            the gradient is normalized by an estimation of its variance
        weight_decay (float, optional): weight decay (L2 penalty) (default: 0)
        decoupled_decay (bool, optional): decoupled weight decay as per https://arxiv.org/abs/1711.05101
        lr_in_momentum (bool, optional): learning rate scaling is included in the momentum buffer
            update as per defaults in Tensorflow

    {Gz??绽|=r           FTc
              
      s   d|kst d|d|ks,t d|d|ksBt d|d|ksXt d|d|ksnt d|t||||||||	d}
tt| ||
 d S )Nr   zInvalid learning rate: {}zInvalid epsilon value: {}zInvalid momentum value: {}zInvalid weight_decay value: {}zInvalid alpha value: {})lrmomentumalphaepscenteredweight_decaydecoupled_decaylr_in_momentum)
ValueErrorformatdictsuperr   __init__)selfparamsr   r
   r   r   r	   r   r   r   defaults	__class__ ^/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/rmsprop_tf.pyr   0   s    zRMSpropTF.__init__c                    s8   t t| | | jD ]}|dd |dd qd S )Nr	   r   r   F)r   r   __setstate__param_groups
setdefault)r   stategroupr   r   r   r   B   s    
zRMSpropTF.__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}|jrttd| j| }t|dkrd|d< t ||d< |d dkrt 	||d< |d	 rt 	||d
< |d }d|d  }|d  d7  < |d dkr:|d r(|
d|d |d    n|j||d d}|j|d| |d |d	 r|d
 }	|	j||	 |d |j|	|	dd|d  }
n||d  }
|d dkr |d }|d r|
|d j||
|d d ||  n*|
|d ||
 |j||d  d qN|j||
|d  d qNq@|S )zPerforms a single optimization step.

        Arguments:
            closure (callable, optional): A closure that reevaluates the model
                and returns the loss.
        Nr   z)RMSprop does not support sparse gradientsr   step
square_avgr	   Zmomentum_bufferr   grad_avgg      ?r
      r   r   r   )r
      )valuer   r   )torchZenable_gradr   gradZ	is_sparseRuntimeErrorr   lenZ	ones_likeZ
zeros_likeZmul_addZadd_powZaddcmulZsqrt_Zaddcdiv_)r   closureZlossr    pr)   r   r"   Zone_minus_alphar#   avgbufr   r   r   r!   H   sR    
$



 
zRMSpropTF.step)r   r   r   r   r   FFT)N)
__name__
__module____qualname____doc__r   r   r(   Zno_gradr!   __classcell__r   r   r   r   r      s   !  r   )r5   r(   Ztorch.optimr   r   r   r   r   r   <module>   s   	