a
    d                     @   s   d Z ddlmZ ddlZddlmZ G dd deZdeej eej eej eee	e	e	e	d	d	d
Z
eej eej eej e	e	e	e	edddZeej eej eej e	e	e	e	edddZdS )z Lion Optimizer
Paper: `Symbolic Discovery of Optimization Algorithms` - https://arxiv.org/abs/2302.06675
Original Impl: https://github.com/google/automl/tree/master/lion
    )ListN)	Optimizerc                       s@   e Zd ZdZd fdd	Z fd	d
Ze dddZ  Z	S )LionzImplements Lion algorithm.-C6?g?gGz?        FNc                    s   d|kst d|d|d   kr.dk sBn t d|d d|d   krZdk snn t d|d t|||||d}t || d	S )
a  Initialize the hyperparameters.

        Args:
          params (iterable): iterable of parameters to optimize or dicts defining
            parameter groups
          lr (float, optional): learning rate (default: 1e-4)
          betas (Tuple[float, float], optional): coefficients used for computing
            running averages of gradient and its square (default: (0.9, 0.99))
          weight_decay (float, optional): weight decay coefficient (default: 0)
        r   zInvalid learning rate: {}r   g      ?z%Invalid beta parameter at index 0: {}   z%Invalid beta parameter at index 1: {})lrbetasweight_decayforeachmaximizeN)
ValueErrorformatdictsuper__init__)selfparamsr	   r
   r   r   r   defaults	__class__ X/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/lion.pyr      s    zLion.__init__c                    s4   t  | | jD ]}|dd |dd  qd S )Nr   Fr   )r   __setstate__param_groups
setdefault)r   stategroupr   r   r   r   ?   s    
zLion.__setstate__c                 C   s  d}|dur:t   | }W d   n1 s00    Y  | jD ]}g }g }g }|d \}}|d D ]n}	|	jdu rtqd||	 |	jjrtd||	j | j|	 }
t|
dkrt j	|	t j
d|
d< ||
d  qdt||||||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.

        Returns:
          the loss.
        Nr
   r   z&Lion does not support sparse gradientsr   )Zmemory_formatexp_avgr	   r   r   r   )beta1beta2r	   r   r   r   )torchZenable_gradr   gradappendZ	is_sparseRuntimeErrorr   lenZ
zeros_likeZpreserve_formatlion)r   closureZlossr   Zparams_with_gradgradsexp_avgsr    r!   pr   r   r   r   stepE   s@    
$



z	Lion.step)r   r   r   FN)N)
__name__
__module____qualname____doc__r   r   r"   Zno_gradr,   __classcell__r   r   r   r   r      s        #r   F)	r   r)   r*   r   r   r    r!   r	   r   c          
   
   C   sV   |du rd}|r"t j r"td|r6t j s6t}	nt}	|	| |||||||d dS )z=Functional API that performs Lion algorithm computation.
    NFz6torch.jit.script not supported with foreach optimizers)r    r!   r	   r   r   )r"   ZjitZis_scriptingr%   _multi_tensor_lion_single_tensor_lion)
r   r)   r*   r   r   r    r!   r	   r   funcr   r   r   r'   z   s"    r'   )r   r)   r*   r    r!   r	   r   r   c                C   s   t | D ]\}}	|s|| n||  }
|| }t|	rVt|
}
t|}t|	}	|	d||   ||j|
d| d}|	jt|| d ||
d|  qd S )Nr   alpha)		enumerater"   
is_complexview_as_realZmul_mulZadd_signZlerp_)r   r)   r*   r    r!   r	   r   r   iparamr#   r   updater   r   r   r3      s    



r3   c          	      C   s   t | dkrd S |r"tt|}dd |D }dd |D }dd | D } t| d||   t||}tj||d| d dd |D }tj| || d t|| tj||d| d d S )	Nr   c                 S   s$   g | ]}t |rt |n|qS r   r"   r8   r9   .0xr   r   r   
<listcomp>       z&_multi_tensor_lion.<locals>.<listcomp>c                 S   s$   g | ]}t |rt |n|qS r   r?   r@   r   r   r   rC      rD   c                 S   s$   g | ]}t |rt |n|qS r   r?   r@   r   r   r   rC      rD   r   r5   c                 S   s   g | ]}|  qS r   )r;   )rA   ur   r   r   rC      rD   )r&   r"   Z_foreach_negtupleZ_foreach_mul_Z_foreach_mulZ_foreach_add_)	r   r)   r*   r    r!   r	   r   r   Zupdatesr   r   r   r2      s    r2   )FN)r0   typingr   r"   Ztorch.optim.optimizerr   r   ZTensorboolfloatr'   r3   r2   r   r   r   r   <module>   sF   g  ) 