
    ^j                     ^    d Z ddlmZ ddlmZ ddlZddlmZmZm	Z	 ddl
mZ  G d d	e      Zy)
an   PyTorch impl of LaProp optimizer

Code simplified from https://github.com/Z-T-WANG/LaProp-Optimizer, MIT License

Paper: LaProp: Separating Momentum and Adaptivity in Adam, https://arxiv.org/abs/2002.04839

@article{ziyin2020laprop,
  title={LaProp: a Better Way to Combine Momentum with Adaptive Gradient},
  author={Ziyin, Liu and Wang, Zhikang T and Ueda, Masahito},
  journal={arXiv preprint arXiv:2002.04839},
  year={2020}
}

References for added functionality:
    Cautious Optimizers: https://arxiv.org/abs/2411.16085
    Why Gradients Rapidly Increase Near the End of Training: https://arxiv.org/abs/2506.02285

    )Tuple)	OptimizerN   )_add_scaled__init_scalar_validate_scalar)ParamsTc                        e Zd ZdZ	 	 	 	 	 	 ddededeeef   dedededef fd	Z fd
Z	 e
j                         dd       Z xZS )LaPropzw LaProp Optimizer

    Paper: LaProp: Separating Momentum and Adaptivity in Adam, https://arxiv.org/abs/2002.04839
    paramslrbetasepsweight_decaycautioncorrected_weight_decayc                 4   t        d|       t        d|       d|d   cxk  rdk  sn t        dj                  |d               d|d   cxk  rdk  sn t        dj                  |d               t        ||||||	      }t        t
        |   ||       y )
Nzlearning rateepsilon        r         ?z%Invalid beta parameter at index 0: {}r   z%Invalid beta parameter at index 1: {})r   r   r   r   r   r   )r   
ValueErrorformatdictsuperr   __init__)
selfr   r   r   r   r   r   r   defaults	__class__s
            \/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/optim/laprop.pyr   zLaProp.__init__!   s     	"-C(eAh$$DKKERSHUVVeAh$$DKKERSHUVV%#9
 	fd$VX6    c                    t         |   |       | j                  D ]  }|j                  dd       |j                  dd       |d   D ]  }| j                  j                  |i       }|s"d|v rt        |d   d      |d<   d|v rFt        j                  |d	         r.t        |d   |d	   j                  |d	   j                  
      |d<   d|v st        |d   d      |d<     y )Nr   Fr   r   stepcpudeviceexp_avg_lr_1r   )dtyper%   exp_avg_lr_2)r   __setstate__param_groups
setdefaultstategetr   torch	is_tensorr'   r%   )r   r,   grouppp_stater   s        r   r)   zLaProp.__setstate__;   s    U#&& 	bEY.5u=8_ b**..B/W$&276?5&QGFO!W,t1M.:/#Dk//$T{11/GN+
 "W,.:7>;R[`.aGN+b	br    c           
         d}|$t        j                         5   |       }ddd       | j                  D ]  }|d   D ]  }|j                  |j                  }|j                  rt        d      | j                  |   }t        |      dk(  rt        d      |d<   t        j                  |      |d<   t        j                  |d	         rt        j                  |d	         nd
|d<   t        d      |d<   t        j                  |      |d<   |d   |d   }}|d   \  }	}
|d   j                  d       d|
z
  }d|	z
  }|j                  |
      j                  |||       |d   |	z  ||d	   z  z   |d<   |d   |
z  |z   |d<   t        j                  |d	         rpt        j                  |d	   d
k7  |d	   t        j                  |d	               }t        j                  |d	   d
k7  |d   |z  t        j                  |d	               }n|d	   d
k7  r|d   |d	   z  nd}|d   }d|z  }|j!                  |      j#                         j                  |d         }||z  }|j                  |	       t%        |||d	   |z         |d   rU||z  dkD  j'                  |j(                        }|j+                  |j-                         j/                  d             ||z  }t%        |||        |d   dk7  s|d   r|d	   dz  | j0                  d	   z  }n|d	   }t%        ||| |d   z           |S # 1 sw Y   xY w)zPerforms a single optimization step.

        Arguments:
            closure (callable, optional): A closure that reevaluates the model
                and returns the loss.
        Nr   z(LaProp does not support sparse gradientsr   r#   r$   r"   exp_avgr   r   r&   r(   
exp_avg_sqr   r   )valuer   r   r   gMbP?)minr   r      )r.   enable_gradr*   grad	is_sparseRuntimeErrorr,   lenr   
zeros_liker/   add_mul_addcmul_where	ones_likedivsqrt_r   tor'   div_meanclamp_r   )r   closurelossr0   r1   r:   r,   r4   r5   beta1beta2one_minus_beta2one_minus_beta1lr_safebias_correction1bias_correction2	step_sizedenomstep_of_this_gradmaskwd_scales                        r   r"   zLaProp.stepO   s    ""$ !y! && A	JE8_ @J66>vv>>&'QRR

1 u:?$0$>E&M','7'7':E)$MR__]bcg]hMiE,<,<U4[,IoqE.),8,FE.)*/*:*:1*=E,'&+I&6l8K$W~uf""1%"#e)"#e) &//d//R(-n(=(EZ_`dZeHe(en%(-n(=(E(Wn% ??5;/#kk%+*;U4[%//Z_`dZeJfgG',{{dr)n-7d4($ OTTXk]_N_u^'<uT{'Jeg$#(#8  00	"'78>>@EEeElS$(5L!U#W&7t9VW##dNQ.224::>DIIdiik00T0:;%nGQ)4(A-56#(;!#3dmmD6I#I#(; Ay53H'HIA@JA	JF M! !s   L99M)g-C6:?)g?g+?gV瞯<r   FF)N)__name__
__module____qualname____doc__r	   floatr   boolr   r)   r.   no_gradr"   __classcell__)r   s   @r   r   r      s     )5"$!+077 7 &	7
 7  7 7 %)74b( U]]_O Or    r   )r[   typingr   torch.optimr   r.   _helpersr   r   r   _typesr	   r    r    r   <module>re      s,   $  !  B B CY Cr    