+
    &jah                        R t ^ RIHt ^ RIt^ RIHt ^RIHtHtHtH	t	H
t
HtHtHtHtHtHtHtHtHtHtHt RR.t ! R R]4      tRR	] R
]
 R] R] R] R2,           ]n         R R ltR R lt]	! ]R7      RR R ll4       tR# )z'Implementation for the NAdam algorithm.)castN)Tensor)_capturable_doc_default_to_fused_or_foreach_differentiable_doc_disable_dynamo_if_unsupported_foreach_doc!_get_capturable_supported_devices_get_scalar_dtype
_get_value_maximize_doc_params_doc_stack_if_compiling
_to_scalar_use_grad_for_differentiable_view_as_real	OptimizerParamsTNAdamnadamc            	       |   a a ] tR t^ t oRRRRRRRRR/V3R lV 3R lllltV 3R	 ltR
 t]RR l4       tRt	Vt
V ;t# )r   FforeachNmaximize
capturabledifferentiablec                   < V ^8  d   QhRS[ RS[S[,          RS[S[S[3,          RS[RS[RS[RS[RS[R	,          R
S[RS[RS[RR	/# )   paramslrbetasepsweight_decaymomentum_decaydecoupled_weight_decayr   Nr   r   r   return)r   floatr   tuplebool)format__classdict__s   "i/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/optim/nadam.py__annotate__NAdam.__annotate__!   s     )+ )+)+ FN)+ UE\"	)+
 )+ )+ )+ !%)+ )+ )+ )+ )+ 
)+    c                 < \        V\        4      '       d!   VP                  4       ^8w  d   \        R4      hRV8:  g   \        RV 24      hRV8:  g   \        RV 24      hRV^ ,          u;8:  d   R8  g   M \        RV^ ,           24      hRV^,          u;8:  d   R8  g   M \        RV^,           24      hRV8:  g   \        RV 24      hRV8:  g   \        R	V 24      hR
VRVRVRVRVRVRV	RVRV
RV/
p\        SV `  W4       R# )   zTensor lr must be 1-element        zInvalid learning rate: zInvalid epsilon value:       ?z#Invalid beta parameter at index 0: z#Invalid beta parameter at index 1: zInvalid weight_decay value: zInvalid momentum_decay value: r   r   r    r!   r"   r#   r   r   r   r   N)
isinstancer   numel
ValueErrorsuper__init__)selfr   r   r   r    r!   r"   r#   r   r   r   r   defaults	__class__s   &&&&&&&&$$$$ r*   r6   NAdam.__init__!   s.    b&!!bhhjAo:;;by6rd;<<cz6se<==eAh$$B58*MNNeAh$$B58*MNNl";L>JKKn$=n=MNOO"U3Ln$&<w*n
 	*r-   c                  < \         SV `  V4       V P                   EF  pVP                  R R4       VP                  RR4       VP                  RR4       VP                  RR4       VP                  RR4       VR,           EFO  pV P                  P                  V. 4      p\        V4      ^ 8w  g   K2  \        P                  ! VR,          4      '       gp   \        VR,          4      pVR,          '       d,   \        P                  ! V\        4       VP                  R	7      M\        P                  ! V\        4       R
7      VR&   \        P                  ! VR,          4      '       d   K  VR,          pVR,          '       d,   \        P                  ! V\        4       VP                  R	7      M\        P                  ! V\        4       R
7      VR&   EKR  	  EK  	  R# )r   Fr   Nr   r   r#   r   stepdtypedevicer>   
mu_product)r5   __setstate__param_groups
setdefaultstategetlentorch	is_tensorr%   tensorr
   r?   )r7   rE   grouppp_statestep_valmu_prod_valr9   s   &&     r*   rB   NAdam.__setstate__L   sr   U#&&EZ/Y-\51-u55u=8__**..B/w<1$ ??76?;;#(#9
  %\22 "LL (0A0CAHH "'h>O>Q!R   !??7<+@AA&-l&;
  %\22 "LL +3D3Fqxx "'kARAT!U  - % 'r-   c                N   R pVR,           EF  p	V	P                   f   K  V\        P                  ! V	4      ,          pVP                  V	4       V	P                   P                  '       d   \        R4      hVP                  V	P                   4       V P                  V	,          p
\        V
4      ^ 8X  Ed   VR,          '       d,   \        P                  ! R\        4       V	P                  R7      M\        P                  ! R\        4       R7      V
R&   VR,          '       d,   \        P                  ! R\        4       V	P                  R7      M\        P                  ! R\        4       R7      V
R	&   \        P                  ! V	\        P                  R
7      V
R&   \        P                  ! V	\        P                  R
7      V
R&   VP                  V
R,          4       VP                  V
R,          4       VP                  V
R	,          4       VP                  V
R,          4       EK  	  V# )Fr   z'NAdam does not support sparse gradientsr   r=   r0   r@   r<   r1   rA   )memory_formatexp_avg
exp_avg_sq )gradrH   
is_complexappend	is_sparseRuntimeErrorrE   rG   zerosr
   r?   rJ   ones
zeros_likepreserve_format)r7   rK   params_with_gradgradsexp_avgsexp_avg_sqsmu_productsstate_stepshas_complexrL   rE   s   &&&&&&&&   r*   _init_groupNAdam._init_groupj   s    xAvv!u//22 ''*66###&'PQQQVV$

1u:? !.. B.?.A!((S"\\#5F5HI &M !.. 

2->-@R"\\#5F5HI ,' (-'7'7)>)>(E)$ +0*:*:)>)>+E,' i 01""5#67""5#67""5=1I !J r-   c                *   V P                  4        RpVe.   \        P                  ! 4       ;_uu_ 4        V! 4       pRRR4       V P                   F  p. p. p. p. p. p. p	\	        \
        \        \        3,          VR,          4      w  rV P                  VVVVVVV	4      p\        VVVVVV	V
VVR,          VR,          VR,          VR,          VR,          VR,          VR	,          VR
,          VR,          VR7       K  	  V#   + '       g   i     L; i)zPerform a single optimization step.

Args:
    closure (Callable, optional): A closure that reevaluates the model
        and returns the loss.
Nr   r   r!   r"   r    r   r#   r   r   r   )beta1beta2r   r!   r"   r    r   r#   r   r   r   re   )	'_accelerator_graph_capture_health_checkrH   enable_gradrC   r   r&   r%   rf   r   )r7   closurelossrK   r_   r`   ra   rb   rc   rd   ri   rj   re   s   &&           r*   r<   
NAdam.step   s'    	446""$$y % &&E-/"$E%'H(*K(*K(*KeUl 3U7^DLE** K  ;">2$%56%Lz*',-E'Fi( .$%56'%' 'P W %$s   DD	rU   )gMb`?)g?g+?g:0yE>    gMbp?FN)__name__
__module____qualname____firstlineno__r6   rB   rf   r   r<   __static_attributes____classdictcell____classcell__)r9   r)   s   @@r*   r   r       s\     )+  $)+ )+ !)+  %)+ )+V<0d "6 "6 6r-   a  Implements NAdam algorithm.

    .. math::
       \begin{aligned}
            &\rule{110mm}{0.4pt}                                                                 \\
            &\textbf{input}      : \gamma_t \text{ (lr)}, \: \beta_1,\beta_2 \text{ (betas)},
                \: \theta_0 \text{ (params)}, \: f(\theta) \text{ (objective)}                   \\
            &\hspace{13mm} \: \lambda \text{ (weight decay)}, \:\psi \text{ (momentum decay)}    \\
            &\hspace{13mm} \: \textit{decoupled\_weight\_decay}, \:\textit{maximize}             \\
            &\textbf{initialize} :  m_0 \leftarrow 0 \text{ ( first moment)},
                v_0 \leftarrow 0 \text{ ( second moment)}                                 \\[-1.ex]
            &\rule{110mm}{0.4pt}                                                                 \\
            &\textbf{for} \: t=1 \: \textbf{to} \: \ldots \: \textbf{do}                         \\
            &\hspace{5mm}\textbf{if} \: \textit{maximize}:                                       \\
            &\hspace{10mm}g_t           \leftarrow   -\nabla_{\theta} f_t (\theta_{t-1})         \\
            &\hspace{5mm}\textbf{else}                                                           \\
            &\hspace{10mm}g_t           \leftarrow   \nabla_{\theta} f_t (\theta_{t-1})          \\
            &\hspace{5mm} \theta_t \leftarrow \theta_{t-1}                                       \\
            &\hspace{5mm} \textbf{if} \: \lambda \neq 0                                          \\
            &\hspace{10mm}\textbf{if} \: \textit{decoupled\_weight\_decay}                       \\
            &\hspace{15mm} \theta_t \leftarrow \theta_{t-1} - \gamma \lambda \theta_{t-1}                    \\
            &\hspace{10mm}\textbf{else}                                                          \\
            &\hspace{15mm} g_t \leftarrow g_t + \lambda \theta_{t-1}                             \\
            &\hspace{5mm} \mu_t \leftarrow \beta_1 \big(1 - \frac{1}{2}  0.96^{t \psi} \big)     \\
            &\hspace{5mm} \mu_{t+1} \leftarrow \beta_1 \big(1 - \frac{1}{2} 0.96^{(t+1)\psi}\big)\\
            &\hspace{5mm}m_t           \leftarrow   \beta_1 m_{t-1} + (1 - \beta_1) g_t          \\
            &\hspace{5mm}v_t           \leftarrow   \beta_2 v_{t-1} + (1-\beta_2) g^2_t          \\
            &\hspace{5mm}\widehat{m_t} \leftarrow \mu_{t+1} m_t/(1-\prod_{i=1}^{t+1}\mu_i)\\[-1.ex]
            & \hspace{11mm} + (1-\mu_t) g_t /(1-\prod_{i=1}^{t} \mu_{i})                         \\
            &\hspace{5mm}\widehat{v_t} \leftarrow   v_t/\big(1-\beta_2^t \big)                   \\
            &\hspace{5mm}\theta_t \leftarrow \theta_t - \gamma \widehat{m_t}/
                \big(\sqrt{\widehat{v_t}} + \epsilon \big)                                       \\
            &\rule{110mm}{0.4pt}                                                          \\[-1.ex]
            &\bf{return} \:  \theta_t                                                     \\[-1.ex]
            &\rule{110mm}{0.4pt}                                                          \\[-1.ex]
       \end{aligned}

    For further details regarding the algorithm we refer to `Incorporating Nesterov Momentum into Adam`_.
    z
    Args:
        a  
        lr (float, Tensor, optional): learning rate (default: 2e-3)
        betas (Tuple[float, float], optional): coefficients used for computing
            running averages of gradient and its square (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)
        momentum_decay (float, optional): momentum momentum_decay (default: 4e-3)
        decoupled_weight_decay (bool, optional): whether to decouple the weight
            decay as in AdamW to obtain NAdamW. If True, the algorithm does not
            accumulate weight decay in the momentum nor variance. (default: False)
        z	
        z

    .. _Incorporating Nesterov Momentum into Adam:
        https://openreview.net/forum?id=OM0jvwB8jIp57ZJjtNEZ
    .. _Decoupled Weight Decay Regularization:
        https://arxiv.org/abs/1711.05101

    c          $      l   V ^8  d   QhR\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\        R\        R	\        R
\        R\        R\        R\        R\        R\        R\        R\        RR/# r   r   r`   ra   rb   rc   rd   ri   rj   r   r!   r"   r    r#   r   r   r   re   r$   Nlistr   r%   r'   )r(   s   "r*   r+   r+     s     a aLa<a 6la f	a
 fa fa a a 	a a a 
a !a a  !a" #a$ %a& 
'ar-   c                   \         P                  P                  4       '       g   \        V4      p\	        V 4       EF  w  ppV'       g
   VV,          M	VV,          ) pVV,          pVV,          pVV,          pVV,          p\         P
                  ! V4      '       dY   \         P                  ! V4      p\         P                  ! V4      p\         P                  ! V4      p\         P                  ! V4      p\         P                  P                  4       '       g   V'       d   \        4       pVP                  P                  VP                  P                  u;8X  d   VP                  P                  8X  d   M MVP                  P                  V9   g   \        R V R24      hV^,          pV'       d   TpM\        V4      p^VV,          ,
          pV	^ 8w  d;   V'       d    VP                  ^W,          ,
          4       MVP                  VV	R7      pVRRRVV
,          ,          ,          ,
          ,          pVRRRV^,           V
,          ,          ,          ,
          ,          pVV,          pVP!                  V^V,
          4       VP                  V4      P#                  VV^V,
          R7       VP%                  V4      P'                  4       pV'       g	   V'       d   VP                  V4      pVV,          pVV) RV,
          ,          RV,
          ,          ,          pVV) V,          RV,
          ,          ,          pVP)                  VV4       VP)                  VV4       EK*  \        V4      V,          pVP+                  V4       VP)                  VVV) RV,
          ,          R\        V4      ,
          ,          R7       VP)                  VV\-        \.        V) V,          RV,
          ,          4      R7       EK  	  R# )zVIf capturable=True, params, mu_products and state_steps must be on supported devices: .alphar1         ?Q?)valueN)rH   jitis_scriptingr   	enumeraterW   view_as_realcompileris_compilingr	   r?   typeAssertionErrorr   mul_addlerp_addcmul_divsqrtaddcdiv_add_r   r%   )r   r`   ra   rb   rc   rd   ri   rj   r   r!   r"   r    r#   r   r   r   re   iparamrV   rS   rT   rA   step_tcapturable_supported_devicesr<   bias_correction2mumu_nextdenommu_product_nexts   &&&&&&$$$$$$$$$$$              r*   _single_tensor_nadamr     s&   ( 99!!##^f%5'uQxeAhY1+ ^
 ^
QE""&&u-E%%d+D((1G++J7J ~~**,,+L+N(!!Z%6%6%;%;Qv}}?Q?QQLL%%)EE$--I,J!M  	!Df%Dud{?1%

1r001xx\x: cC4D>,A#BCCD3$(n1L(M!NNO 	b
 	dAI&''d!e)'D/0557ZIIcNE )72OB3#(+sZ/?@AD"w#2G!HIGNN4'NN7E*(4w>OJJsONNeRC38$4j>T8T$U   NN5B3=S?5J"KL  M &r-   c          $      l   V ^8  d   QhR\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\        R\        R	\        R
\        R\        R\        R\        R\        R\        R\        R\        RR/# rz   r{   )r(   s   "r*   r+   r+   }  s     \ \L\<\ 6l\ f	\
 f\ f\ \ \ 	\ \ \ 
\ !\ \  !\" #\$ %\& 
'\r-   c                v  a, \        V 4      ^ 8X  d   R# V'       d   \        R4      h\        P                  P	                  4       '       g|   V'       dt   \        RR7      o,\        ;QJ d+    V,3R l\        WVRR7       4       F  '       d   K   RM 	  RM! V,3R l\        WVRR7       4       4      '       g   \        RS, R	24      h\        V4      p\        P                  ! WW#WE.4      pVP                  4        EF  w  w  ppppppp\        \        \        ,          V4      p\        \        \        ,          V4      p\        \        \        ,          V4      p\        \        \        ,          V4      p\        \        \        ,          V4      p\        \        \        ,          V4      pV'       d   \        VVVV4       V'       d   \        P                   ! V4      p\        P                  P	                  4       '       gJ   V^ ,          P"                  '       d1   \        P$                  ! V\        P&                  ! R
RR7      R
R7       M\        P$                  ! V^4       V	^ 8w  di   V'       d&   \        P(                  ! V^W,          ,
          4       M;V'       d   \        P$                  ! VVV	R7       M\        P*                  ! VVV	R7      p\        P,                  ! VV^V,
          4       \        P(                  ! VV4       \        P.                  ! VVV^V,
          4       \        P0                  ! V4      pV'       EdC   \        P2                  ! VV
4      p \        P4                  ! RV 4      p!\        P(                  ! V!R4       \        P$                  ! V!R
4       \        P(                  ! V!V4       \        P$                  ! V V
4       \        P4                  ! RV 4      p"\        P(                  ! V"R4       \        P$                  ! V"R
4       \        P(                  ! V"V4       ? \        P4                  ! VV4      p#\        P6                  ! V#R
4       \        P8                  ! V#4       \        P:                  ! V#4       MV U$u. uF#  p$^V\=        V$4      ,          ,
          R,          NK%  	  p#p$V U$u. uF1  p$VR
RR\=        V$4      V
,          ,          ,          ,
          ,          NK3  	  p!p$V U$u. uF8  p$VR
RR\=        V$4      ^,           V
,          ,          ,          ,
          ,          NK:  	  p"p$\        P(                  ! VV!4       \        P>                  ! VV#4       \        P$                  ! VV4       ?#V'       Ed   \        P6                  ! V!R
4       \        P(                  ! V!V4       \        P@                  ! VR
4      p%\        P8                  ! V%4       \        P>                  ! V!V%4       T!p&?%\        P2                  ! VV"4      p%\        P(                  ! V"V4       \        P6                  ! V%R
4       \        P>                  ! V"V%4       T"p'?%\        P2                  ! V&V4      p(\        P.                  ! V(V'V4       \        PB                  ! VV(V4       EK  \E        \        VV!RR7       U)U*u. uF=  w  p)p*\=        V4      R
V*,
          ,          R
\=        V)4      ,
          ,          R,          NK?  	  up*p)4      p&\E        \        VV"RR7       U)U+u. uF=  w  p)p+\=        V4      V+,          R
\=        V)4      V+,          ,
          ,          R,          NK?  	  up+p)4      p'\        PB                  ! VVVV&4       \        PB                  ! VVVV'4       EK  	  R# u up$i u up$i u up$i u up*p)i u up+p)i )rp   Nz#_foreach ops don't support autogradF)supports_xlac              3     <"   T F{  w  rpVP                   P                  VP                   P                  u;8H  ;'       d    VP                   P                  8H  Mu ;'       d    VP                   P                  S9   x  K}  	  R # 5irq   )r?   r   ).0rL   mpr<   r   s   &   r*   	<genexpr>&_multi_tensor_nadam.<locals>.<genexpr>  sf      
  Rt HHMMRYY^^??t{{/?/?? > >!==>Qs   A B$"BT)strictzWIf capturable=True, params, mu_products, and state_steps must be on supported devices: r~   r1   cpu)r?   r   r   r   g      )#rG   r   rH   r   r   r	   allzipr   r   "_group_tensors_by_device_and_dtypevaluesr   r|   r   r   _foreach_negis_cpu_foreach_add_rJ   _foreach_mul__foreach_add_foreach_lerp__foreach_addcmul__foreach_sqrt_foreach_mul_foreach_pow_foreach_sub__foreach_neg__foreach_sqrt_r   _foreach_div__foreach_sub_foreach_addcdiv_r   )-r   r`   ra   rb   rc   rd   ri   rj   r   r!   r"   r    r#   r   r   r   re   grouped_tensorsgrouped_params_grouped_grads_grouped_exp_avgs_grouped_exp_avg_sqs_grouped_mu_products_grouped_state_steps__grouped_paramsgrouped_gradsgrouped_exp_avgsgrouped_exp_avg_sqsgrouped_mu_productsgrouped_state_stepsexp_avg_sq_sqrtexponentmusmu_nextsbias_correction_sqrtr<   r   step_size_gradsstep_size_expavg	numeratorrA   r   r   r   s-   &&&&&&$$$$$$$$$$$                           @r*   _multi_tensor_nadamr   }  s   ( 6{aBCC >>&&((Z'H(
$ s 
  #6DQ
sss 
  #6DQ
 
 

 !V/03  
BBBB	{HO ""$		 	d6lO<T&\>:V.?@"4<1EF"4<1EF"4<1EF /?AT !..}=M ~~**,,1DQ1G1N1N1N#U\\#e%DC  3Q71%##NA8I4IJ ''%~\ %*$6$6%~\%M
 	-}a%iH/7q5y	
  --.AB
 :))*=~NH$$T84CT*S)U+ .9))$9H$/#.%0 #(#5#5e=P#Q  4c: 45  !56 DW$CV4Uj...366CV ! $
 0/D sdz$/?./P&QRRSS/   0 0D *T*:Q*>.)P QRRT T/   	/5O-ABOS1 !:S)R(&&':C@E&U+!O &&':HEE"- s+%0' **?MJI##I/?AQR ##NIO1 +..A3t*T*T
B  ^sRx0C*Z:P4PQUWWW*TO  3 03+Xd0
0+
G #2!"J!7'!AAC  0
  ##	 ##  	C %`$b
s    )^ ;7^%8>^*;A^/A^5)single_tensor_fnc          &         V ^8  d   QhR\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\        R\        R	,          R
\        R\        R\        R\        R\        R\        R\        R\        R\        R\        RR	/# )r   r   r`   ra   rb   rc   rd   r#   r   Nr   r   re   r   ri   rj   r   r!   r"   r    r$   )r|   r   r'   r%   )r(   s   "r*   r+   r+   ]  s     D DLD<D 6lD f	D
 fD fD !D D[D D D D D  !D" #D$ 	%D& 'D( )D* 
+D, 
-Dr-   c               T   \         ;QJ d    R V 4       F  '       d   K   RM	  RM! R V 4       4      '       g   \        R4      h\         ;QJ d    R V 4       F  '       d   K   RM	  RM! R V 4       4      '       g   \        R4      hVf   \        W	RR7      w  ppV'       d0   \        P                  P                  4       '       d   \        R	4      hV'       d,   \        P                  P                  4       '       g   \        pM\        pV! V VVVVVVVVVVVVVVV	V
R
7       R# )zhFunctional API that performs NAdam algorithm computation.

See :class:`~torch.optim.NAdam` for details.
c              3   V   "   T F  p\        V\        P                  4      x  K!  	  R # 5irq   r2   rH   r   r   ts   & r*   r   nadam.<locals>.<genexpr>x       @Kqz!U\\**K   ')FTzPAPI has changed, `state_steps` argument must contain a list of singleton tensorsc              3   V   "   T F  p\        V\        P                  4      x  K!  	  R # 5irq   r   r   s   & r*   r   r   }  r   r   zPAPI has changed, `mu_products` argument must contain a list of singleton tensorsN)	use_fusedz6torch.jit.script not supported with foreach optimizers)ri   rj   r   r!   r"   r   r#   r    r   r   re   )r   rZ   r   rH   r   r   r   r   )r   r`   ra   rb   rc   rd   r#   r   r   r   re   r   ri   rj   r   r!   r"   r    r   funcs   &&&&&&&&&&&&$$$$$$  r*   r   r   \  s    8 3@K@333@K@@@^
 	
 3@K@333@K@@@^
 	
 1e

7 599))++STTuyy--//"#!%5%#r-   )FNFFFF)__doc__typingr   rH   r   	optimizerr   r   r   r   r   r	   r
   r   r   r   r   r   r   r   r   r   __all__r   r   r   r   rU   r-   r*   <module>r      s    .       ( G
sI sn&N		 	 
 		 		 		 !O> FaH\~  1EFD GDr-   