a
    d7                     @   s8  d dl Z d dl mZ ddlmZmZmZmZmZmZ d dl	m
Z
mZ d dlmZ ddgZG d	d deZd
djeeed e_de
e e
e e
e e
e e
e ee eeeeeeeedddZe
e e
e e
e e
e e
e eeeeeeeedddZe
e e
e e
e e
e e
e eeeeeeeedddZdS )    N)Tensor   )	Optimizer_default_to_fused_or_foreach_use_grad_for_differentiable_differentiable_doc_foreach_doc_maximize_doc)ListOptional)"_group_tensors_by_device_and_dtypeRMSproprmspropc                	       sN   e Zd Zdee eed fdd	Z fd
dZdd ZedddZ	  Z
S )r   {Gz?Gz?:0yE>r   FNforeachmaximizedifferentiablec                    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 || d S )Ng        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_decayr   r   r   )
ValueErrorformatdictsuper__init__)selfparamsr   r   r   r   r   r   r   r   r   defaults	__class__ \/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/optim/rmsprop.pyr       s,    zRMSprop.__init__c                    sX   t  | | jD ]@}|dd |dd |dd  |dd |dd qd S )Nr   r   r   Fr   r   r   )r   __setstate__param_groups
setdefault)r!   stategroupr$   r&   r'   r(   1   s    
zRMSprop.__setstate__c           	      C   s$  |d D ]}|j d u rq|| |j jr4td||j  | j| }t|dkrd|d< tj|tjd|d< |d dkrtj|tjd|d< |d	 rtj|tjd|d
< ||d  |d dkr||d  |d	 r||d
  |d rt	|d t
rtd|d  d7  < qd S )Nr"   z)RMSprop does not support sparse gradientsr   step)Zmemory_format
square_avgr   Zmomentum_bufferr   grad_avgr   z`step` can't be a tensorr   )gradappendZ	is_sparseRuntimeErrorr+   lentorchZ
zeros_likeZpreserve_format
isinstancer   )	r!   r,   params_with_gradgradssquare_avgsmomentum_buffer_list	grad_avgspr+   r&   r&   r'   _init_group:   s:    





zRMSprop._init_groupc           	      C   s   d}|dur:t   | }W d   n1 s00    Y  | jD ]t}g }g }g }g }g }| |||||| t||||||d |d |d |d |d |d |d |d	 |d
 d q@|S )zPerforms a single optimization step.

        Args:
            closure (Callable, optional): A closure that reevaluates the model
                and returns the loss.
        Nr   r   r   r   r   r   r   r   r   )	r   r   r   r   r   r   r   r   r   )r4   Zenable_gradr)   r<   r   )	r!   closureZlossr,   r6   r7   r8   r:   r9   r&   r&   r'   r-   `   s8    
$
zRMSprop.step)	r   r   r   r   r   FNFF)N)__name__
__module____qualname__r   boolr    r(   r<   r   r-   __classcell__r&   r&   r$   r'   r      s$            	%	&a  Implements RMSprop algorithm.

    .. math::
       \begin{aligned}
            &\rule{110mm}{0.4pt}                                                                 \\
            &\textbf{input}      : \alpha \text{ (alpha)},\: \gamma \text{ (lr)},
                \: \theta_0 \text{ (params)}, \: f(\theta) \text{ (objective)}                   \\
            &\hspace{13mm}   \lambda \text{ (weight decay)},\: \mu \text{ (momentum)},\: centered\\
            &\textbf{initialize} : v_0 \leftarrow 0 \text{ (square average)}, \:
                \textbf{b}_0 \leftarrow 0 \text{ (buffer)}, \: g^{ave}_0 \leftarrow 0     \\[-1.ex]
            &\rule{110mm}{0.4pt}                                                                 \\
            &\textbf{for} \: t=1 \: \textbf{to} \: \ldots \: \textbf{do}                         \\
            &\hspace{5mm}g_t           \leftarrow   \nabla_{\theta} f_t (\theta_{t-1})           \\
            &\hspace{5mm}if \: \lambda \neq 0                                                    \\
            &\hspace{10mm} g_t \leftarrow g_t + \lambda  \theta_{t-1}                            \\
            &\hspace{5mm}v_t           \leftarrow   \alpha v_{t-1} + (1 - \alpha) g^2_t
                \hspace{8mm}                                                                     \\
            &\hspace{5mm} \tilde{v_t} \leftarrow v_t                                             \\
            &\hspace{5mm}if \: centered                                                          \\
            &\hspace{10mm} g^{ave}_t \leftarrow g^{ave}_{t-1} \alpha + (1-\alpha) g_t            \\
            &\hspace{10mm} \tilde{v_t} \leftarrow \tilde{v_t} -  \big(g^{ave}_{t} \big)^2        \\
            &\hspace{5mm}if \: \mu > 0                                                           \\
            &\hspace{10mm} \textbf{b}_t\leftarrow \mu \textbf{b}_{t-1} +
                g_t/ \big(\sqrt{\tilde{v_t}} +  \epsilon \big)                                   \\
            &\hspace{10mm} \theta_t \leftarrow \theta_{t-1} - \gamma \textbf{b}_t                \\
            &\hspace{5mm} else                                                                   \\
            &\hspace{10mm}\theta_t      \leftarrow   \theta_{t-1} -
                \gamma  g_t/ \big(\sqrt{\tilde{v_t}} + \epsilon \big)  \hspace{3mm}              \\
            &\rule{110mm}{0.4pt}                                                          \\[-1.ex]
            &\bf{return} \:  \theta_t                                                     \\[-1.ex]
            &\rule{110mm}{0.4pt}                                                          \\[-1.ex]
       \end{aligned}

    For further details regarding the algorithm we refer to
    `lecture notes <https://www.cs.toronto.edu/~tijmen/csc321/slides/lecture_slides_lec6.pdf>`_ by G. Hinton.
    and centered version `Generating Sequences
    With Recurrent Neural Networks <https://arxiv.org/pdf/1308.0850v5.pdf>`_.
    The implementation here takes the square root of the gradient average before
    adding epsilon (note that TensorFlow interchanges these two operations). The effective
    learning rate is thus :math:`\gamma/(\sqrt{v} + \epsilon)` where :math:`\gamma`
    is the scheduled learning rate and :math:`v` is the weighted moving average
    of the squared gradient.
    a  
    Args:
        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 constant (default: 0.99)
        eps (float, optional): term added to the denominator to improve
            numerical stability (default: 1e-8)
        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)
        {foreach}
        {maximize}
        {differentiable}

    r   F)r"   r7   r8   r:   r9   r   r   r   r   r   r   r   r   r   c                C   sn   |du rt | |dd\}}|r0tj r0td|rDtj sDt}nt}|| ||||||	|
|||||d dS )zsFunctional API that performs rmsprop algorithm computation.
    See :class:`~torch.optim.RMSProp` for details.
    NF)Z	use_fusedz6torch.jit.script not supported with foreach optimizers)r   r   r   r   r   r   r   r   )r   r4   ZjitZis_scriptingr2   _multi_tensor_rmsprop_single_tensor_rmsprop)r"   r7   r8   r:   r9   r   r   r   r   r   r   r   r   r   _funcr&   r&   r'   r      s,    )r"   r7   r8   r:   r9   r   r   r   r   r   r   r   r   c                C   sT  t | D ]D\}}|| }|s"|n| }|| }|dkrF|j||d}t|}|rrt|}t|}t|}||j||d| d |
r|| }|rt|}||j|d| d |j||dd	 }n|
 }|r||}n
||}|	dkr<|| }|rt|}||	|| |j|| d q|j||| d qd S )Nr   r   r   value)	enumerateaddr4   
is_complexview_as_realZmul_Zaddcmul_Zadd_ZaddcmulZsqrt_sqrtZaddcdiv_)r"   r7   r8   r:   r9   r   r   r   r   r   r   r   r   iparamr0   r.   Zis_complex_paramr/   avgbufr&   r&   r'   rD      s:    







rD   c                C   sv  t | dkrd S |rJ dt| ||||g}| D ]8\}}}}}|rTt|}|dkrltj|||d}dd }||}||}||}t|| tj|||d| d |
r||}t|| tj||d| d tj	|||dd}t
| t|| nt|}t|| |	dkr\||}t||	 t||| tj||| d q6tj|||| d q6d S )	Nr   z#_foreach ops don't support autogradrG   c                 S   s   dd | D S )Nc                 S   s$   g | ]}t |rt |n|qS r&   )r4   rM   rN   ).0tr&   r&   r'   
<listcomp>S  s   zH_multi_tensor_rmsprop.<locals>._view_complex_as_real.<locals>.<listcomp>r&   )Ztensor_listr&   r&   r'   _view_complex_as_realR  s    z4_multi_tensor_rmsprop.<locals>._view_complex_as_realr   rH   rJ   )r3   r   valuesr4   Z_foreach_negZ_foreach_addZ_foreach_mul_Z_foreach_addcmul_Z_foreach_add_Z_foreach_addcmulZ_foreach_sqrt_Z_foreach_sqrtZ_foreach_addcdiv_)r"   r7   r8   r:   r9   r   r   r   r   r   r   r   r   Zgrouped_tensorsZgrouped_paramsZgrouped_gradsZgrouped_square_avgsZgrouped_grad_avgsZgrouped_momentum_buffer_listrW   rR   r&   r&   r'   rC   3  s@    



rC   )NFF)r4   r   Z	optimizerr   r   r   r   r   r	   typingr
   r   Ztorch.utils._foreach_utilsr   __all__r   r   __doc__rA   floatr   rD   rC   r&   r&   r&   r'   <module>   st    *E   4: