
    ^j                     \    d Z ddlZddlmc mZ ddlmZmZ ddl	Z	ddl
mZ  G d de      Zy)a  
SGDP Optimizer Implementation copied from https://github.com/clovaai/AdamP/blob/master/adamp/sgdp.py

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

References for added functionality:
    Cautious Optimizers: https://arxiv.org/abs/2411.16085
    Spherical Cautious Optimizers: https://openreview.net/forum?id=OyT2CJ4fh7 
Copyright (c) 2020-present NAVER Corp.
MIT license
    N)	Optimizerrequired   )
projectionc            	       b     e Zd Zeddddddddf	 fd	Z ej                         dd       Z xZS )SGDPr   Fg:0yE>g?c                 V    t        ||||||||	|
	      }t        t        |   ||       y )N)	lrmomentum	dampeningweight_decaynesterovepsdeltawd_ratiocaution)dictsuperr   __init__)selfparamsr
   r   r   r   r   r   r   r   r   defaults	__class__s               Z/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/optim/sgdp.pyr   zSGDP.__init__   s=     %

 	dD"684    c                    d }|$t        j                         5   |       }d d d        | j                  D ]  }|d   }|d   }|d   }|d   }|j                  dd      }|d   D ]s  }	|	j                  |	j                  }
| j
                  |	   }t        |      dk(  rt        j                  |	      |d<   |d   }|j                  |      j                  |
d	|z
  
       |r	|
||z  z   }n|j                         }d	}t        |	j                        dkD  rt        |	|
||d   |d   |d   |      \  }}nc|ra||
z  dkD  j                  |
j                        }|j                  |j!                         j#                  d             |j                  |       |dk7  r&|	j                  d	|d   |d   z  |z  d|z
  z  z
         |	j                  ||d    
       v  |S # 1 sw Y   xY w)Nr   r   r   r   r   Fr   r   g      ?)alphar   r   r   r   gMbP?)minr
   )torchenable_gradparam_groupsgetgradstatelen
zeros_likemul_add_cloneshaper   todtypediv_meanclamp_)r   closurelossgroupr   r   r   r   r   pr#   r$   bufd_pr   masks                   r   stepz	SGDP.step1   s   ""$ !y! && '	0E 0LZ(Hk*IZ(Hii	51G8_  066>vv

1 u:?(-(8(8(;E*% J'"''BN'CC/C))+C qww<!#$.q$U7^US]M^`efk`lnu$vMC$JN..tzz:DIIdiik00T0:;HHTN  1$FF2deN.C Ch NRST\R\ ]]^ s5;,/A 0'	0R Y! !s   G  G*)N)	__name__
__module____qualname__r   r   r   no_gradr7   __classcell__)r   s   @r   r   r      sA     54 U]]_/ /r   r   )__doc__r   torch.nn.functionalnn
functionalFtorch.optim.optimizerr   r   mathadampr   r    r   r   <module>rF      s,       5  K9 Kr   