a
    d                     @   sx   d Z ddlZddlm  mZ ddlmZ ddlZej	dddZ
ej	dddZeeed	d
dZG dd deZdS )a3  
AdamP Optimizer Implementation copied from https://github.com/clovaai/AdamP/blob/master/adamp/adamp.py

Paper: `Slowing Down the Weight Norm Increase in Momentum-based Optimizers` - https://arxiv.org/abs/2006.08217
Code: https://github.com/clovaai/AdamP

Copyright (c) 2020-present NAVER Corp.
MIT license
    N)	Optimizer)returnc                 C   s   |  | ddS )Nr   )reshapesizex r	   Y/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/adamp.py_channel_view   s    r   c                 C   s   |  ddS )N   r   )r   r   r	   r	   r
   _layer_view   s    r   )deltawd_ratioepsc                 C   s   d}ddt | jd   }ttfD ]}|| }	||}
tj|
|	d|d }| |t	|	
d k r"| |	jddd|| }||||| jdd| 8 }|}||f  S q"||fS )	N      ?)r   )r   r   )dimr      )pr   )r   )lenshaper   r   FZcosine_similarityZabs_maxmathsqrtr   Znormadd_r   sum)r   gradperturbr   r   r   wdZexpand_sizeZ	view_funcZ
param_viewZ	grad_viewZ
cosine_simZp_nr	   r	   r
   
projection   s    "r    c                       s0   e Zd Zd fdd	Ze dd
dZ  ZS )AdamPMbP?g?g+?:0yE>r   皙?Fc	           
   	      s,   t |||||||d}	tt| ||	 d S )N)lrbetasr   weight_decayr   r   nesterov)dictsuperr!   __init__)
selfparamsr&   r'   r   r(   r   r   r)   defaults	__class__r	   r
   r,   ,   s
    zAdamP.__init__Nc              
   C   s  d }|d ur:t   | }W d    n1 s00    Y  | jD ]}|d D ]}|jd u r`qN|j}|d \}}|d }| j| }	t|	dkrd|	d< t ||	d< t ||	d< |	d |	d  }
}|	d  d7  < d||	d   }d||	d   }|
|j|d| d	 ||j	||d| d
 |
 t
| |d }|d | }|rp||
 d| |  | }n|
| }d}t|jdkrt||||d |d |d \}}|d dkr|d|d |d  |   |j|| d	 qNq@|S )Nr.   r'   r)   r   stepexp_avg
exp_avg_sqr   )alpha)valuer   r&   r   r   r   r(   )torchZenable_gradZparam_groupsr   stater   Z
zeros_likeZmul_r   Zaddcmul_r   r   r   r    )r-   closureZlossgroupr   r   Zbeta1Zbeta2r)   r8   r3   r4   Zbias_correction1Zbias_correction2ZdenomZ	step_sizer   r   r	   r	   r
   r2   3   sD    
$

"z
AdamP.step)r"   r#   r$   r   r%   r%   F)N)__name__
__module____qualname__r,   r7   Zno_gradr2   __classcell__r	   r	   r0   r
   r!   +   s
     r!   )__doc__r7   Ztorch.nn.functionalnnZ
functionalr   Ztorch.optim.optimizerr   r   ZTensorr   r   floatr    r!   r	   r	   r	   r
   <module>   s   
