a
    þd‡  ã                   @   s$   d Z ddlZG dd„ dejjƒZdS )z› AdaHessian Optimizer

Lifted from https://github.com/davda54/ada-hessian/blob/master/ada_hessian.py
Originally licensed MIT, Copyright 2020, David Samuel
é    Nc                       s`   e Zd ZdZd‡ fd	d
„	Zedd„ ƒZdd„ Zdd„ Ze	 
¡ dd„ ƒZe	 
¡ ddd„ƒZ‡  ZS )Ú
Adahessianaî  
    Implements the AdaHessian algorithm from "ADAHESSIAN: An Adaptive Second OrderOptimizer for Machine Learning"

    Arguments:
        params (iterable): iterable of parameters to optimize or dicts defining parameter groups
        lr (float, optional): learning rate (default: 0.1)
        betas ((float, float), optional): coefficients used for computing running averages of gradient and the
            squared hessian trace (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 (L2 penalty) (default: 0.0)
        hessian_power (float, optional): exponent of the hessian trace (default: 1.0)
        update_each (int, optional): compute the hessian trace approximation only after *this* number of steps
            (to save time) (default: 1)
        n_samples (int, optional): how many times to sample `z` for the approximation of the hessian trace (default: 1)
    çš™™™™™¹?©gÍÌÌÌÌÌì?g+‡ÙÎ÷ï?ç:Œ0âŽyE>ç        ç      ð?é   Fc
                    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 s„n t d|d › ƒ‚d|  kr˜dks¨n t d	|› ƒ‚|| _|| _|	| _d
| _t ¡  | j¡| _t	|||||d}
t
t| ƒ ||
¡ |  ¡ D ]}d|_d| j| d< qþd S )Nr   zInvalid learning rate: zInvalid epsilon value: r   r   z#Invalid beta parameter at index 0: r   z#Invalid beta parameter at index 1: zInvalid Hessian power value: iÿÿÿ)ÚlrÚbetasÚepsÚweight_decayÚhessian_powerúhessian step)Ú
ValueErrorÚ	n_samplesÚupdate_eachÚavg_conv_kernelÚseedÚtorchÚ	GeneratorÚmanual_seedÚ	generatorÚdictÚsuperr   Ú__init__Ú
get_paramsÚhessÚstate)ÚselfÚparamsr	   r
   r   r   r   r   r   r   ÚdefaultsÚp©Ú	__class__© ú^/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/adahessian.pyr      s(    zAdahessian.__init__c                 C   s   dS )NTr$   ©r   r$   r$   r%   Úis_second_order6   s    zAdahessian.is_second_orderc                 C   s   dd„ | j D ƒS )zH
        Gets all parameters in all param_groups with gradients
        c                 s   s&   | ]}|d  D ]}|j r|V  qqdS )r   N)Zrequires_grad)Ú.0Úgroupr!   r$   r$   r%   Ú	<genexpr>?   ó    z(Adahessian.get_params.<locals>.<genexpr>)Úparam_groupsr&   r$   r$   r%   r   :   s    zAdahessian.get_paramsc                 C   s@   |   ¡ D ]2}t|jtƒs| j| d | j dkr|j ¡  qdS )z;
        Zeros out the accumalated hessian traces.
        r   r   N)r   Ú
isinstancer   Úfloatr   r   Zzero_)r   r!   r$   r$   r%   Úzero_hessianA   s    $zAdahessian.zero_hessianc           	   	      s  g }t dd„ ˆ  ¡ ƒD ]<}ˆ j| d ˆ j dkr<| |¡ ˆ j| d  d7  < qt|ƒdkrddS ˆ jj|d jkrt 	|d j¡ 
ˆ j¡ˆ _dd„ |D ƒ}tˆ jƒD ]f}‡ fd	d„|D ƒ}tjj|||d
|ˆ jd k d}t|||ƒD ]"\}}}| j|| ˆ j 7  _qêq¨dS )z}
        Computes the Hutchinson approximation of the hessian trace and accumulates it for each trainable parameter.
        c                 S   s
   | j d uS )N©Úgrad)r!   r$   r$   r%   Ú<lambda>Q   r+   z(Adahessian.set_hessian.<locals>.<lambda>r   r   r   Nc                 S   s   g | ]
}|j ‘qS r$   r0   ©r(   r!   r$   r$   r%   Ú
<listcomp>\   r+   z*Adahessian.set_hessian.<locals>.<listcomp>c              	      s0   g | ](}t jd d| ¡ ˆ j|jdd d ‘qS )r   é   )r   Údeviceg       @r   )r   ÚrandintÚsizer   r6   r3   r&   r$   r%   r4   `   r+   T)Zgrad_outputsZonly_inputsZretain_graph)Úfilterr   r   r   ÚappendÚlenr   r6   r   r   r   r   Úranger   Zautogradr1   Úzipr   )	r   r   r!   ZgradsÚiZzsZh_zsZh_zÚzr$   r&   r%   Úset_hessianJ   s"    
ÿzAdahessian.set_hessianNc                 C   s°  d}|dur|ƒ }|   ¡  |  ¡  | jD ]€}|d D ]p}|jdu s6|jdu rRq6| jrŒ| ¡ dkrŒt |j¡j	ddgdd 
|j¡ ¡ |_| d|d	 |d
   ¡ | j| }t|ƒdkràd|d< t |¡|d< t |¡|d< |d |d  }}|d \}}	|d  d7  < | |¡j|jd| d | |	¡j|j|jd|	 d d||d   }
d|	|d   }|d }||  |d ¡ |d ¡}|d	 |
 }|j||| d q6q(|S )z¿
        Performs a single optimization step.
        Arguments:
            closure (callable, optional) -- a closure that reevaluates the model and returns the loss (default: None)
        Nr   é   r5   é   T)ÚdimZkeepdimr   r	   r   r   ÚstepÚexp_avgÚexp_hessian_diag_sqr
   )Úalpha)Úvaluer   r   )r/   r@   r,   r1   r   r   rC   r   ÚabsÚmeanZ	expand_asÚcloneZmul_r   r;   Z
zeros_likeZadd_Zaddcmul_Zpow_Zaddcdiv_)r   ÚclosureZlossr)   r!   r   rE   rF   Zbeta1Zbeta2Zbias_correction1Zbias_correction2ÚkZdenomZ	step_sizer$   r$   r%   rD   f   s:    (
zAdahessian.step)r   r   r   r   r   r   r   F)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úpropertyr'   r   r/   r   Zno_gradr@   rD   Ú__classcell__r$   r$   r"   r%   r   	   s     ÿ
	
r   )rQ   r   ZoptimZ	Optimizerr   r$   r$   r$   r%   Ú<module>   s   