a
    d}c                     @   s  d dl Z d dl mZ ddlmZmZmZmZmZ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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 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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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eeeeeeeddddZdS )    N)Tensor   )	Optimizer_use_grad_for_differentiable
_get_value_dispatch_sqrt_stack_if_compiling_capturable_doc_differentiable_doc_foreach_doc
_fused_doc_maximize_doc_default_to_fused_or_foreach)ListOptional)"_group_tensors_by_device_and_dtypeAdamWadamwc                       sd   e Zd Zdddddddeee eeee d fdd	Z fd
dZdd ZedddZ	  Z
S )r   MbP?g?g+?:0yE>{Gz?FN)maximizeforeach
capturabledifferentiablefusedc                   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 sn t d|d d|kst d	|t||||||||	|
|d

}t || |r|
rtdd| _tdd | jD std|rtdd S )N        zInvalid learning rate: {}zInvalid epsilon value: {}r   g      ?z%Invalid beta parameter at index 0: {}r   z%Invalid beta parameter at index 1: {}zInvalid weight_decay value: {})
lrbetasepsweight_decayamsgradr   r   r   r   r   z)`fused` does not support `differentiable`Tc                 s   s,   | ]$}|d  D ]}|j o t|V  qqdS )paramsN)is_cudatorchZis_floating_point).0Zpgp r(   Z/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/optim/adamw.py	<genexpr><   s   z!AdamW.__init__.<locals>.<genexpr>zF`fused=True` requires all the params to be CUDA, floating point Tensorz0`fused` and `foreach` cannot be `True` together.)	
ValueErrorformatdictsuper__init__RuntimeErrorZ_step_supports_amp_scalingallparam_groups)selfr#   r   r   r    r!   r"   r   r   r   r   r   defaults	__class__r(   r)   r/      sB    zAdamW.__init__c                    s   t  | | jD ]L}|dd |dd |dd  |dd |dd |dd  qt| j }t|dkot	|d d	 }|s|D ]}t
t|d	 |d	< qd S )
Nr"   Fr   r   r   r   r   r   step)r.   __setstate__r2   
setdefaultliststatevalueslenr%   Z	is_tensortensorfloat)r3   r;   groupZstate_valuesZstep_is_tensorsr5   r(   r)   r8   D   s    

zAdamW.__setstate__c	                 C   s  |d D ]}	|	j d u rq||	 |	j jr2td||	j  | j|	 }
t|
dkr|d sd|d rxtjdtj|	j	dnt
d|
d	< tj|	tjd
|
d< tj|	tjd
|
d< |rtj|	tjd
|
d< ||
d  ||
d  |r||
d  ||
d	  qd S )Nr#   z'AdamW does not support sparse gradientsr   r   r   r   )dtypedevicer   r7   )Zmemory_formatexp_avg
exp_avg_sqZmax_exp_avg_sq)gradappendZ	is_sparser0   r;   r=   r%   zerosr?   rD   r>   Z
zeros_likeZpreserve_format)r3   r@   params_with_gradgradsr"   exp_avgsexp_avg_sqsmax_exp_avg_sqsstate_stepsr'   r;   r(   r(   r)   _init_groupU   s<    





zAdamW._init_groupc                 C   s   |    d}|durBt  | }W d   n1 s80    Y  | jD ]}g }g }g }g }g }g }	|d }
|d \}}| ||||
||||	 t||||||	|
|||d |d |d |d |d |d	 |d
 |d t| ddt| ddd qH|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   
grad_scale	found_inf)r"   beta1beta2r   r!   r    r   r   r   r   r   rQ   rR   )Z _cuda_graph_capture_health_checkr%   Zenable_gradr2   rP   r   getattr)r3   closureZlossr@   rJ   rK   rL   rM   rN   rO   r"   rS   rT   r(   r(   r)   r7      s\    
$


z
AdamW.step)r   r   r   r   F)N)__name__
__module____qualname__boolr   r/   r8   rP   r   r7   __classcell__r(   r(   r5   r)   r      s,        		72a  Implements AdamW algorithm.

    .. math::
       \begin{aligned}
            &\rule{110mm}{0.4pt}                                                                 \\
            &\textbf{input}      : \gamma \text{(lr)}, \: \beta_1, \beta_2
                \text{(betas)}, \: \theta_0 \text{(params)}, \: f(\theta) \text{(objective)},
                \: \epsilon \text{ (epsilon)}                                                    \\
            &\hspace{13mm}      \lambda \text{(weight decay)},  \: \textit{amsgrad},
                \: \textit{maximize}                                                             \\
            &\textbf{initialize} : m_0 \leftarrow 0 \text{ (first moment)}, v_0 \leftarrow 0
                \text{ ( second moment)}, \: \widehat{v_0}^{max}\leftarrow 0              \\[-1.ex]
            &\rule{110mm}{0.4pt}                                                                 \\
            &\textbf{for} \: t=1 \: \textbf{to} \: \ldots \: \textbf{do}                         \\

            &\hspace{5mm}\textbf{if} \: \textit{maximize}:                                       \\
            &\hspace{10mm}g_t           \leftarrow   -\nabla_{\theta} f_t (\theta_{t-1})          \\
            &\hspace{5mm}\textbf{else}                                                           \\
            &\hspace{10mm}g_t           \leftarrow   \nabla_{\theta} f_t (\theta_{t-1})           \\
            &\hspace{5mm} \theta_t \leftarrow \theta_{t-1} - \gamma \lambda \theta_{t-1}         \\
            &\hspace{5mm}m_t           \leftarrow   \beta_1 m_{t-1} + (1 - \beta_1) g_t          \\
            &\hspace{5mm}v_t           \leftarrow   \beta_2 v_{t-1} + (1-\beta_2) g^2_t          \\
            &\hspace{5mm}\widehat{m_t} \leftarrow   m_t/\big(1-\beta_1^t \big)                   \\
            &\hspace{5mm}\widehat{v_t} \leftarrow   v_t/\big(1-\beta_2^t \big)                   \\
            &\hspace{5mm}\textbf{if} \: amsgrad                                                  \\
            &\hspace{10mm}\widehat{v_t}^{max} \leftarrow \mathrm{max}(\widehat{v_t}^{max},
                \widehat{v_t})                                                                   \\
            &\hspace{10mm}\theta_t \leftarrow \theta_t - \gamma \widehat{m_t}/
                \big(\sqrt{\widehat{v_t}^{max}} + \epsilon \big)                                 \\
            &\hspace{5mm}\textbf{else}                                                           \\
            &\hspace{10mm}\theta_t \leftarrow \theta_t - \gamma \widehat{m_t}/
                \big(\sqrt{\widehat{v_t}} + \epsilon \big)                                       \\
            &\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 `Decoupled Weight Decay Regularization`_.
    a  
    Args:
        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-8)
        weight_decay (float, optional): weight decay coefficient (default: 1e-2)
        amsgrad (bool, optional): whether to use the AMSGrad variant of this
            algorithm from the paper `On the Convergence of Adam and Beyond`_
            (default: False)
        {maximize}
        {foreach}
        {capturable}
        {differentiable}
        {fused}
    .. _Decoupled Weight Decay Regularization:
        https://arxiv.org/abs/1711.05101
    .. _On the Convergence of Adam and Beyond:
        https://openreview.net/forum?id=ryQu7f-RZ

    )r   r   r   r   r   F)r#   rK   rL   rM   rN   rO   r   r   r   r   rQ   rR   r"   rS   rT   r   r!   r    r   c                C   s   t dd |D std|	du r<|du r<t| |dd\}}|	du rHd}	|du rTd}|rjtj rjtd|	rtj rtd|	rtj st}n|rtj st}nt}|| |||||||||||||||
|d	 dS )
zpFunctional API that performs AdamW algorithm computation.

    See :class:`~torch.optim.AdamW` for details.
    c                 s   s   | ]}t |tjV  qd S N)
isinstancer%   r   )r&   tr(   r(   r)   r*   %      zadamw.<locals>.<genexpr>zPAPI has changed, `state_steps` argument must contain a list of singleton tensorsNF)Z	use_fusedz6torch.jit.script not supported with foreach optimizersz4torch.jit.script not supported with fused optimizers)r"   rS   rT   r   r!   r    r   r   r   rQ   rR   )	r1   r0   r   r%   ZjitZis_scripting_fused_adamw_multi_tensor_adamw_single_tensor_adamw)r#   rK   rL   rM   rN   rO   r   r   r   r   rQ   rR   r"   rS   rT   r   r!   r    r   _funcr(   r(   r)   r     sL    )r#   rK   rL   rM   rN   rO   rQ   rR   r"   rS   rT   r   r!   r    r   r   r   c       	         C   s@  |d u r|d u sJ t | D ]\}}|s2|| n||  }|| }|| }|| }|rl|jrd|jslJ dt|rt|}t|}t|}t|}|d7 }|d||   ||	j|d|	 d ||
j||d|
 d |s|r|}dt|	| }dt|
| }|| }|	 }|
 }|r|rJ||  }n|| }|| t|| || 
 ||  || }n|
 ||  || }||| qt|}d|	|  }d|
|  }|| }t|}|rtj|| ||| d || 
 | |}n|
 | |}|j||| d qd S )N@If capturable=True, params and state_steps must be CUDA tensors.r   alpha)value)out)	enumerater$   r%   
is_complexview_as_realZmul_Zadd_Zaddcmul_pownegsqrtcloneZcopy_maximumZaddcdiv_r   r   )r#   rK   rL   rM   rN   rO   rQ   rR   r"   rS   rT   r   r!   r    r   r   r   iparamrG   rE   rF   Zstep_tr7   bias_correction1bias_correction2	step_sizeZstep_size_negbias_correction2_sqrtZmax_exp_avg_sqs_idenomr(   r(   r)   rb   V  sj    





rb   c       	             s  t | dkrd S |r4tdd t| |D s4J d|r@J d|d u rP|d u sTJ t| |||||g}| D ]\}}}}}}|rtt|}dd |D }dd |D }d	d |D }d
d |D }t|d t	|d|   t	|  tj||d  d t	| t
|||d  |rP fdd|D }fdd|D }t|d t|d t| t| t|}t| t| t|}|r t|| t|}t|t|| t||}t| t||}n@t|}t|t|| t||}t| t||}t||| qp fdd|D }fdd|D }tfdd|D }dd |D }|rt|| t|}t|| t||}n"t|}t|| t||}t|||| qpd S )Nr   c                 s   s   | ]\}}|j o|j V  qd S r\   )r$   )r&   r'   r7   r(   r(   r)   r*     s   z&_multi_tensor_adamw.<locals>.<genexpr>re   z#_foreach ops don't support autogradc                 S   s$   g | ]}t |rt |n|qS r(   r%   rk   rl   r&   xr(   r(   r)   
<listcomp>  r_   z'_multi_tensor_adamw.<locals>.<listcomp>c                 S   s$   g | ]}t |rt |n|qS r(   ry   rz   r(   r(   r)   r|     r_   c                 S   s$   g | ]}t |rt |n|qS r(   ry   rz   r(   r(   r)   r|     s   c                 S   s$   g | ]}t |rt |n|qS r(   ry   rz   r(   r(   r)   r|     r_   r   rf   c                    s   g | ]}t  |qS r(   r%   rm   r&   r7   rS   r(   r)   r|     r_   c                    s   g | ]}t  |qS r(   r}   r~   rT   r(   r)   r|     r_   c                    s   g | ]}d  t |  qS rB   r   r~   r   r(   r)   r|   $  r_   c                    s   g | ]}d  t |  qS rB   r   r~   r   r(   r)   r|   %  r_   c                    s   g | ]} | d  qS )r(   r&   Zbc)r   r(   r)   r|   '  r_   c                 S   s   g | ]}t |qS r(   )r   r   r(   r(   r)   r|   )  r_   )r=   r1   zipr   r<   r%   Z_foreach_negtuple_foreach_add_Z_foreach_mul_Z_foreach_addcmul__foreach_sub_Z_foreach_neg_Z_foreach_divZ_foreach_reciprocal_Z_foreach_sqrtZ_foreach_maximum_Z_foreach_div_Z_foreach_mulZ_foreach_addZ_foreach_addcdiv_r   ) r#   rK   rL   rM   rN   rO   rQ   rR   r"   rS   rT   r   r!   r    r   r   r   grouped_tensorsdevice_paramsdevice_gradsdevice_exp_avgsdevice_exp_avg_sqsdevice_max_exp_avg_sqsdevice_state_stepsrt   ru   rv   rw   Zmax_exp_avg_sq_sqrtZeps_over_step_sizerx   Zexp_avg_sq_sqrtr(   )rS   rT   r   r)   ra     s    












ra   )r#   rK   rL   rM   rN   rO   rQ   rR   r"   rS   rT   r   r!   r    r   r   r   returnc       	         C   s$  |rt d|d ur|j|ind }|d ur4|j|ind }t| |||||g}|D ]\}}|||f \}}}}}}d\}}|d ur||vr|j|dd||< || }|d ur||vr|j|dd||< || }t|d tj|||||||||	|
|||||d |d urPt||gt|  qPd S )Nz"_fused_adamw is not differentiable)NNT)Znon_blockingr   )	r"   r   rS   rT   r!   r    r   rQ   rR   )	r0   rD   r   tor%   r   Z_fused_adamw_r   r=   )r#   rK   rL   rM   rN   rO   rQ   rR   r"   rS   rT   r   r!   r    r   r   r   Zgrad_scale_dictZfound_inf_dictr   rD   rC   r   r   r   r   r   r   Zdevice_grad_scaleZdevice_found_infr(   r(   r)   r`   ;  sV    
r`   )NFFNNN)r%   r   Z	optimizerr   r   r   r   r   r	   r
   r   r   r   r   typingr   r   Ztorch.utils._foreach_utilsr   __all__r   r,   __doc__rZ   r?   r   rb   ra   r`   r(   r(   r(   r)   <module>   s   4 9&M      Oh