
    ^j.3              '       4   d Z ddlmZmZmZmZ ddlZddlmZ ddlm	Z	 ddl
mZmZmZ ddlmZ d	eed
f   dededeeeef      fdZ G d de	      Zdee   dee   deee      deee      deee      deee      dee   dededededededee   deeej.                  f   dee   d ed!ed"ee   f&d#Zdee   dee   deee      deee      deee      deee      dee   dededededededee   deeej.                  f   dee   d ed!ed"ee   f&d$Zy)%a   Adafactor (Big Vision variant) for PyTorch

Adapted from the implementation in big vision: https://github.com/google-research/big_vision

Described in 'Scaling Vision Transformers': https://arxiv.org/abs/2106.04560

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

Adaptation and PyTorch modifications by Ross Wightman
    )ListOptionalTupleUnionN)Tensor)	Optimizer   )
_get_value_init_scalar_validate_scalar)ParamsTshape.factoredmin_dim_size_to_factorreturnc                     |rt        |       dk  ryt        d t        |       D              }| |d   d      |k  ryt        |d   d         t        |d   d         fS )a  Whether to use a factored second moment estimator.

    This function returns a tuple with the two largest axes to reduce over.
    If no two dimensions have size >= min_dim_size_to_factor, return None.

    Args:
      shape: an input shape
      factored: whether to use factored second-moment estimator for > 2d vars.
      min_dim_size_to_factor: only factor accumulator if two array dimensions have at least this size.

    Returns:
      None or a tuple of ints
       Nc              3   *   K   | ]  \  }}||f  y wN ).0ixs      b/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/optim/adafactor_bv.py	<genexpr>z!_factored_dims.<locals>.<genexpr>+   s     >TQ1a&>s   r	   )lensorted	enumerateint)r   r   r   sorted_dimss       r   _factored_dimsr#      sj    $ s5zA~>Yu-=>?K[_Q #99{2q!"CB(:$;;;    c            !           e Zd ZdZddddddej
                  dd	dd
d
d
fd
ddedededededede	e   de
eej                  f   de	e   dede	e   dededede	e   f fdZ fdZ ej                          dd       Z xZS )AdafactorBigVisionz
    PyTorch implementation of BigVision's Adafactor variant with both single and multi tensor implementations.

    Adapted from https://github.com/google-research/big_vision by Ross Wightman
          ?   g?r   g+?g?Ng        F)foreachparamslrr   
decay_ratedecay_offset	beta2_capmomentummomentum_dtypeepsweight_decayclipping_thresholdunscaled_wdcautioncorrected_weight_decayr)   c                >   t        |t              rK|dk(  rt        j                  }n5|dk(  rt        j                  }n|dk(  s
J | d       t        j
                  }t        d|       t        d|
       t        ||||||||	|
|||||      }t        | %  ||       y )Nfloat16bfloat16float32z dtype not supportedzlearning rater2   )r+   r   r,   r-   r.   r/   r0   r1   r2   r3   r4   r5   r6   r)   )

isinstancestrtorchr8   r9   r:   r   dictsuper__init__)selfr*   r+   r   r,   r-   r.   r/   r0   r1   r2   r3   r4   r5   r6   r)   defaults	__class__s                    r   r@   zAdafactorBigVision.__init__8   s    & nc**!&:-!&%2[~6FFZ4[[2!&"-6#9!%)%1##9
  	*r$   c                    t         |   |       | j                  D ]  }|j                  dd       |j                  dd       |j                  dd        |d   D ]  }| j                  j                  |i       }t        |      dk7  r&d|v r"t        |d   dt        j                  	      |d<   d
|v sXt        j                  |d
         sq|d
   j                  | j                  d         |d
<     y )Nr5   Fr6   r)   r*   r   stepcpudevicedtypeexp_avgr0   rI   )r?   __setstate__param_groups
setdefaultstategetr   r   r=   float64	is_tensortorB   )rA   rO   grouppp_staterC   s        r   rL   zAdafactorBigVision.__setstate__i   s    U#&& 	fEY.5u=Y-8_ 	f**..B/w<1$7):&276?5X]XeXe&fGFO'EOOGI<N,O *1);)>)>T]]ScEd)>)eGI&	f		fr$   c                    d }|$t        j                         5   |       }d d d        | j                  D ]   }g }g }g }g }g }g }	g }
|d   D ]`  }|j                  |j                  j                  rt        d      |j                  |       |j                  |j                         | j                  |   }t        |      dk(  rMt        dt         j                        |d<   |j                  j                  }t        |d| j                  d   	      }||\  }}t        |j                  j                        }d
||<   t        |j                  j                        }d
||<   |j                  j                  |      |d<   |j                  j                  |      |d<   n2t        j                   |j                  t         j"                        |d<   | j                  d   1t        j                   |j                  | j                  d         |d<   |	j                  |d          |j                  |j%                  dd              |j                  |j%                  dd              |j                  |j%                  dd              |
j                  |j%                  dd              c |d   rt&        }nt(        } |d%i d|d|d|d|d|d|
d|	d|d   d|d   d|d   d|d   d|d   d|d   d|d   d|d   d |d    d!|d!   d"|d"   d#|d$   r| j                  d   nd   |S # 1 sw Y   xY w)&Nr*   zSparse gradients not supportedr   rF   rG   rE   Tr   )r   r   r	   exp_avg_sq_rexp_avg_sq_c)memory_format
exp_avg_sqr/   r0   rK   rJ   r)   gradsexp_avg_sq_rsexp_avg_sq_csexp_avg_sqsexp_avgsstate_stepsbeta2_decayr,   r.   r1   r+   r2   r3   r4   r5   max_lrr6   r   )r=   enable_gradrM   grad	is_sparseRuntimeErrorappendrO   r   r   rQ   r   r#   rB   list	new_zeros
zeros_likepreserve_formatrP   _multi_tensor_adafactor_single_tensor_adafactor)rA   closurelossrT   params_with_gradr\   r]   r^   r_   ra   r`   rU   rO   r   factored_dimsdcdr	row_shape	col_shapefuncs                       r   rE   zAdafactorBigVision.stepz   s   ""$ !y! && L	E!EMMKKH8_ (<66>66##&'GHH ''*QVV$

1u:?$0U]]$SE&MFFLLE$2!%/3}}=U/V%M %0!.B$($6	()	"$($6	()	"010@0@0Kn-010@0@0Kn-.3.>.>qvvUZUjUj.kl+}}Z0<+0+;+;AFF$--XhJi+ji(""5=1$$UYY~t%DE$$UYY~t%DE""599\4#@A		)T :;Q(<T Y./ ' , ,	
 ( " ( ",/  , (--E'F %L ; #>2 z*  %%56  $))=#>!" "-0#$ i(%& /44L.Mt}}T*SW'qL	\ c! !s   L;;Mr   )__name__
__module____qualname____doc__r=   r9   r   floatr!   r   r   r<   rI   boolr@   rL   no_gradrE   __classcell__)rC   s   @r   r&   r&   1   s    *, # !$(+6;nn#'"%26 %!+0/+" ',#/+/+ /+ %(	/+
 /+ /+ /+ uo/+ "#u{{"23/+ %/+  /+ !)/+ /+ /+ %)/+" d^#/+bf" U]]_T Tr$   r&   r*   r\   r]   r^   r_   r`   ra   rb   r.   r1   r+   r2   r/   r0   r3   r4   r5   rc   c          	      N   t        |       D ]  \  }}||   }||   }||   }||   }||   }||   }|
!|j                  t        j                  k(  rdnd}
|j	                  d       t        |      }t        j                  |      r[t        j                  ||j                  |j                        }t        j                  |dt        j                  ||       z
        }nt        |d|| z  z
        }d|z
  }t        j                  |      |
z   }|t        |j                  d|	      \  } }!|j                  |j!                  |!d      |       |j                  |j!                  | d      |       | |!kD  r| dz
  n| }"|j!                  |"d      }#||#z  j#                         }$|j#                         }%||$z  |%z  }&n+||J |j                  ||       ||j#                         z  }&|I|&j%                  d	      |&j'                         d
z  |z  z  j)                  d      }'|&j+                  |'       ||||j                  k7  r@|j                  |&j-                  |      d|z
         |j-                  |j                        }&n%|j                  |&d|z
         |j/                         }&|ra|&|z  dkD  j-                  |j                        }(|(j+                  |(j!                         j)                  d             |&j1                  |(       |&j1                  |       |dk7  rk|r2||j1                  d|z
         nR|j1                  d||z  |z  z
         n7||j1                  d||z  z
         n|j1                  d|d	z  |z  |z  z
         |j	                  |&d        y )NgHz>gKH9r	   )rI   rH   r'   T)r   )dimkeepdimr   g      ?)maxr   gMbP?)ming      )alpha)r    rI   r=   r8   add_r
   rR   tensorrH   minimumpowr   squarer#   r   lerp_meanrsqrtnormnumelclamp_div_rS   clonemul_))r*   r\   r]   r^   r_   r`   ra   rb   r.   r   r1   r+   r2   r/   r0   r3   r4   r5   rc   r   paramre   rX   rY   r[   rJ   step_trE   beta2_cap_tbeta2_tone_minus_beta2_tgrad_sqrrs   rt   	reduce_dcrow_col_mean
row_factor
col_factorupdatedenommasks)                                            r   rn   rn      s}   , f% R'5Qx$Q'$Q' ^
1+Q;**5$5C 	A&!??4 ,,y

4;;WKmmKuyy|7T1TUG)S4[L+A%ABGK<<%+#DJJMcdFBx}}T}BDUVx}}T}BDUV"$r'QrI',,D,IL&5<<>J%++-JJ&3F  'L,@@@X'89J,,..F )[[^#(=AS'ST\\ad\eEKK G$7+fii7XF DJJ/fa(l3 )--djj9		$))+,,,67D! 	B 1>JJrL01 JJrR&[L$@@A >JJrB$556 JJrR1Wv%5$EEF 	

6
&eR'r$   c                    J d       )Nz2multi-tensor fn (foreach=True) not implemented yetr   )r*   r\   r]   r^   r_   r`   ra   rb   r.   r   r1   r+   r2   r/   r0   r3   r4   r5   rc   s                      r   rm   rm   =  s    . GFF5r$   )r{   typingr   r   r   r   r=   r   torch.optimr   _helpersr
   r   r   _typesr   r!   r}   tupler#   r&   r|   r<   rI   rn   rm   r   r$   r   <module>r      s   0 /   ! @ @ <S#X<< !$< eCHo	<4^ ^Bh'Vh'F|h' HV,-h' HV,-	h'
 (6*+h' x'(h' &\h' h' h' !$h' h' h' h' 5/h'  c5;;./!h'" %UO#h'$ %h'& 'h'( )h'VGVGF|G HV,-G HV,-	G
 (6*+G x'(G &\G G G !$G G G G 5/G  c5;;./!G" %UO#G$ %G& 'G( )Gr$   