
    ^j              )          d Z ddlZddlZddlZddlmZmZmZmZm	Z	m
Z
 ddlZ	 ddlmZ dZddlmZmZmZmZmZ dd	lmZ dd
lmZ ddlmZ  ej4                  e      ZdZdZdgg dg dg ddZe
e e	e!e!e!f   ee	e!e!e!f      f   Z"dFdejF                  de$dejF                  fdZ%de!de	e&df   de!fdZ'edejP                  dfdejF                  de&dee	e!e!e!f      de!de!d ejR                  d!e$dejF                  fd"Z*	 dGd#ejV                  d$e de!fd%Z,	 dGd#ejV                  d$e de	e!e$f   fd&Z-	 	 	 dHdejF                  d'e&d(e!d)e$de
e$e	e$e f   f   f
d*Z.	 dId+ejF                  d,e de	ejF                  ejV                  f   fd-Z/dd.d/eejF                     d0eejF                     d1eejF                     d2e!d3e!d4e!d5e$d6e&d7e"de!de!d$ee    d8e d9e$d!e$de$ddf"d:Z0dd.d/eejF                     d0eejF                     d1eejF                     d;eejF                     d<eejF                     d2e!d3e!d4e!d5e$d=e!d6e&d7e"de!de!d$ee    d8e d9e$d!e$de$ddf(d>Z1dd.d/eejF                     d0eejF                     d1eejF                     d2e!d3e!d4e!d5e$d6e&d7e"de!de!d$ee    d8e d9e$d!e$de$ddf"d?Z2dd.d/eejF                     d0eejF                     d1eejF                     d;eejF                     d<eejF                     d2e!d3e!d4e!d5e$d=e!d6e&d7e"de!de!d$ee    d8e d9e$d!e$de$ddf(d@Z3 G dA dBejh                  jj                        Z6dCe
e ee!   eee!      f   dDee eee!      f   dee	e!e!e!f      fdEZ7y# e$ r dZY <w xY w)Ja   Muon Optimizer

Improved Muon optimizer implementation with flexible handling of high-dimensional tensors.

Combines PyTorch-style structure with options for:
- Batched spatial processing for convolutions in addition to flatten
- Optional spatial normalization
- Selectable coefficient presets
- Automatic fallback to AdamW for 1D / scalar parameters (biases, norms, etc.) and optional fallback via param groups
- AdaMuon (https://arxiv.org/abs/2507.11005)
- mUP eps damping factor (https://arxiv.org/abs/2512.05620v1)

TODO look into mUP LR scaling and independent weight-decay scale

Based on implementation by Keller Jordan, see
- https://github.com/KellerJordan/Muon/blob/master/muon.py
- https://github.com/KellerJordan/modded-nanogpt/blob/master/train_gpt.py
- https://github.com/KellerJordan/modded-nanogpt/blob/master/train_gpt_medium.py
- https://github.com/NoahAmsel/PolarExpress/blob/main/polar_express.py

Hacked together by Ross Wightman
    N)ListMappingOptionalSequenceTupleUnion)DTensorTF   )_add_scaled__check_capturable_devices
_get_value_validate_scalar_zeros_scalar)ParamsT)adamw)nadamwgHz>   )guV@ggn@ @))gͪV@gg"~j@)gv@ggj+6gF%u@)ga4@gH}]g\Cm@)g2%@g?$	g/L
F?)g6>W[@gQkgH}8?))g8y @gWr"b(7g=90@)g3kT@g)6$g}U"?)g")i@gV}gT"?)g߸%~b
@g_O"Dveg)E?)gT`@gU4N/g?)gͶ?g!̲gG߆?)g_?g]sig??N!?)g9+_?gWg-|?))gh('P @g־{6g]!/@)g=.+@gixgےl% ?)g""@gP+.gPYJخ1?)g.69I
@gG4gH5?)gSO|g@gVg@ę?)gΐX?gxDgPh?)gkei?gB?Ggd݈ؼ?)g1?gjgZ?)originalquinticpolar_expresspolar_express_saferparam
capturablereturnc                 @    t        |r| j                        S d       S )N)device)r   r   )r   r   s     Z/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/optim/muon.py
_init_stepr   Z   s    
EEEE    epsshape.c                 ,    |d   |d   }}| ||z  dz  z  S )u  Scale epsilon for Newton-Schulz based on matrix dimensions (μP-style).

    For μP compatibility, epsilon should scale as eps * sqrt(din/dout) to maintain
    consistent damping behavior across different model widths.

    Reference: https://arxiv.org/abs/2512.05620

    Args:
        eps: Base epsilon value
        shape: Shape of the matrix (out, in) or (batch, out, in)

    Returns:
        Scaled epsilon value
          ? )r    r!   doutdins       r   scale_eps_for_nsr)   ^   s)    ( rE"I#D#*$$$r         ?Gstepscoefficientssafety_factordtype	scale_epsc           	      2   | j                   dv sJ d| j                    d       t        |      }|dk\  rt        |d         dk(  sJ ||k  r|d| n||d   g||z
  z  z   }|rt        || j                        }| j	                  |d	
      }	|	j                  d      |	j                  d      kD  }
|
r|	j                  }	|rB|	j                  |	j                  ddd	      j                  |      j                  |             nB|	j                  |	j                  ddd	      j                  |      j                  |             t        xr t        | t              }|r4|D ].  \  }}}|	|	j                  z  }||z  |||z  z  z   }||	z  ||	z  z   }	0 n|	j                   dkD  rt        j                   nt        j"                  }|	j%                         }	t        j&                  g |	j                  dd |	j                  d      |	j(                  |	j*                        }t        j,                  |      }t        j,                  |	      }|D ]>  \  }}} |||	|	j                  dd|        |||||||        ||	||	|d|       ||	}}	@ |
r|	j                  }	|	S )u  Newton-Schulz quintic iteration to compute the zeroth power / orthogonalization of gradient.

    Supports batched operation over leading dimensions.

    See
    - https://github.com/KellerJordan/Muon/blob/master/muon.py
    - https://github.com/NoahAmsel/PolarExpress/blob/main/polar_express.py
    - https://github.com/KellerJordan/modded-nanogpt/blob/master/train_gpt.py

    Args:
        G: Input gradient tensor of shape (m, n) or (batch, m, n)
        steps: Number of Newton-Schulz iterations
        coefficients: Coefficients (a, b, c) for the iteration
        eps: Numerical stability epsilon for norm
        safety_factor: Multiplicative safety factor for norm (1.01 is common safety value in 'polar express' variants)
        dtype: Computation dtype
        scale_eps: If True, scale epsilon by sqrt(din/dout) for μP compatibility

    Returns:
        Orthogonalized tensor of same shape as G
    )      zInput must be 2D or 3D, got zD. Flatten batch dims first.r
   r   r3   Nr$   T)r/   copyr#   r2   )r#   r$   )dimkeepdim)min)r   r/   g        r*   )betaalphaout)ndimlenr)   r!   tosizemTdiv_normmuladd_clamp_has_dtensor
isinstancer	   torchbaddbmmaddmm
contiguousemptyr   r/   
empty_like)r+   r,   r-   r    r.   r/   r0   num_cscoeff_sequenceX
transposed
is_dtensorabcABmm_fnCs                      r   zeropower_via_newtonschulzrY   v   s   < 66V`;AFF8C_``FQ;3|A/1444-2f_\&5)R()UV^<<  sAGG,	5t$A affRj(JDD 	qvvaXtv488GLLSQR	qvvaXtv488GNNSVNWX7Aw!7J% 	 GAq!ADDAAQU#AAQA	  "#! LLNKK3!''#2,3r
3AHHAGGTQQ & 	GAq!!Q3cq9!Q2!Q!4aqA		 DDHr   param_shapeadjust_lr_fnc                     t        |       dkD  r
| d   | d   fnd\  }}|dk(  rt        d||z        dz  S |dk(  rdt        ||      dz  z  S |d	k(  r||z  dz  S J d
| d       )av  Adjust learning rate based on parameter shape for Muon.

    Args:
        param_shape: Shape of the parameter tensor
        adjust_lr_fn: Scaling function name
            - "original": sqrt(max(1, out/in)) - Original Muon impl
            - "match_rms_adamw": 0.2 * sqrt(max(out, in)) - Kimi scaling
            - "rms_to_rms": sqrt(out/in) - Scion/Bernstein scaling
    r
   r#   r$   r*   r*   r   r%   match_rms_adamw皙?
rms_to_rmsInvalid scaling function "z
" for Muon)r<   maxrZ   r[   out_chsin_chss       r   get_lr_scalerf      s     =@<Lq<P{2B8V^OGVz!1g&'3..	*	*S&)S000		% & S((K2<.
KKur   c                     t        |       dkD  r
| d   | d   fnd\  }}|dk(  rd||z  dz  z  dfS |d	k(  r
||z  dz  d
fS |dk(  r|dz  d
fS J d| d       )zAdjust learning rate based on parameter shape for AdaMuon.

    Args:
        param_shape: Shape of the parameter tensor
        adjust_lr_fn: Scaling function name

    Returns:
        Tuple of (scale_factor, use_rms_norm)
    r
   r#   r$   r]   r^   r_   r%   Tr`   Frsqrt_in      ra   z" for AdaMuon)r<   rc   s       r   get_adamuon_lr_scalerj      s     =@<Lq<P{2B8V^OGV(( g&3..44		%& S(%//		#~u$$N2<.NNur   min_dim_sizemax_aspect_ratioreturn_reasonc                 r   | j                   }| j                  dk  st        d |D              dk  r|rdS dS |d   dk(  s|d   dk(  r|rdS dS | j                  dk\  r|d   }d}|dd	 D ]  }||z  }	 ||f}n|}t        |      }	|	|k  r
|rdd
|	 fS yt	        |      }
|
|	z  }||kD  r|rdd|dfS y|rdS dS )a  Check if a parameter is suitable for Muon optimization.

    Args:
        param: Parameter tensor
        min_dim_size: Minimum size for non-unit dimensions
        max_aspect_ratio: Maximum allowed aspect ratio
        return_reason: If True, return (bool, reason_string), else just bool (faster)

    Returns:
        If return_reason=False: bool indicating suitability
        If return_reason=True: Tuple of (is_suitable, reason_string)

    Examples:
        (64, 128) -> True (or (True, "ok") if return_reason=True)
        (96, 3, 4, 4) -> True - will be flattened to (96, 48)
        (4, 2048) -> False - extreme aspect ratio
        (64,) -> False - insufficient dims
        (1, 196, 768) -> False - leading unit dims

    NOTE: these rules were created to balance complexity with covering common timm model cases
    Please let me know if there are non-optimal cases that you run into.
    r2   c              3   ,   K   | ]  }|d kD  s	d   yw)r
   Nr&   ).0dim_sizes     r   	<genexpr>z(_is_suitable_for_muon.<locals>.<genexpr>$  s     A8HqLQAs   
)Finsufficient_dimsFr   r
   )Fleading_unit_dimsr3   Nzmin_dim_too_small:zextreme_aspect_ratio:z.1f)TokT)r!   r;   sumr7   rb   )r   rk   rl   rm   sout_chin_ch_with_spatiald
check_dimsmin_sizemax_sizeaspect_ratios               r   _is_suitable_for_muonr     s   : 	AzzA~AqAAAE/<+G%G 	tqyAaDAI/<+G%GzzQ 112 	$A!#	$01
 
 :H,.xj999 :Hh&L&&1,s1CDDD(<2d2r   tensormodec                 ^   | j                   }| j                  dk(  r| |fS | j                  dk  rt        d| j                         | j                   dd \  }}|dk(  r| j                  |d      |fS |dk(  r*| j                  ||d      }|j	                  ddd      }||fS t        d	|       )
a  Reshape high-dimensional tensor for Muon processing.

    Args:
        tensor: Input tensor of shape (out, in, *spatial)
        mode: How to handle spatial dimensions
            - "flatten": Flatten spatial into output dimension (out, in*H*W)
            - "batched": Batch over spatial positions (spatial_prod, out, in) for per-position orthogonalization

    Returns:
        Reshaped tensor and original shape for restoration
    r2   z,Tensor must have at least 2 dimensions, got Nflattenr$   batchedr   r
   zUnknown mode: )r!   r;   
ValueErrorreshapepermute)r   r   original_shaperx   in_chreshapeds         r   reshape_for_muonr   M  s     \\N{{a~%%{{QG}UVVLL!$MFEy~~fb)>99		 >>&%4##Aq!,''>$011r   )r   paramsgradsmomentum_bufslrweight_decaymomentumnesterovns_stepsns_coefficients	conv_modenormalize_spatialc                :    t        | |||||||||	|
|||||       y)z8Functional API that performs Muon algorithm computation.r   r   r   r   r   r   r    r.   r[   r   r   r0   r   N)_single_tensor_muon)r   r   r   r   r   r   r   r   r   r    r.   r[   r   r   r0   r   s                   r   muonr   p  s<    ( !'#!+!r   exp_avg_sqsstate_stepsbeta2c                >    t        | ||||f|||||	|
||||||||d y)a<  Functional API that performs AdaMuon algorithm computation.

    AdaMuon extends Muon with element-wise second moment estimation applied
    to orthogonalized update directions, providing Adam-like adaptive scaling
    while preserving Muon's geometric benefits.

    Reference: https://arxiv.org/abs/2507.11005
    r   r   r   r   r   r   r   r    r.   r[   r   r   r0   r   N)_single_tensor_adamuon)r   r   r   r   r   r   r   r   r   r   r   r   r    r.   r[   r   r   r0   r   s                      r   adamuonr     sL    <  !'#!+'r   c          	      v   t        |t              }t        |       D ]  \  }}||   }||   }|j                  d||z  z
         |j	                  |d|z
         |r|j	                  ||      n|j                         }|j                  dk\  rt        ||      \  }}n|}|j                  }t        ||||	|
|      }|rt        |j                  |      }nd}|dk(  r9|j                  dk\  r*|r||j                  d   dz  z  }|j                  dd	d      }|j                  |      }t        ||| |z          y
)zSingle tensor Muon update.r
   r*   r3   r   r    r.   r0   r   r   ri   r2   N)resolve_ns_coefficients_COEFFICIENTS	enumeratemul_lerp_cloner;   r   r!   rY   rf   r   r   r   )r   r   r   r   r   r   r   r   r   r    r.   r[   r   r   r0   r   ir   gradmomentum_bufupdateupdate_reshapedr   update_orthoscales                            r   r   r     s\   ( .o}MOf% -75Qx$Q' 	

1rL(() 	4h/7?L(3\EWEWEY ;;!.>vI.V+O^$O#\\N 2'
  !3!3\BEE 	!l&7&71&< ++A.$66'//1a8L $++N; 	UL2#+6[-7r   c          	         t        |t              }|rt        | |       t        |       D ],  \  }}||   }||   }||   }||   }|j	                  d       |r|n
t        |      }|j                  d||z  z
         |j                  |d|z
         |r|j                  ||      n|j                         }|j                  dk\  rt        ||      \  }}n|}|j                  }t        ||
||||      }|dk(  r"|j                  dk\  r|j                  ddd      }|j                  |      }|j                  |	      j                  ||d|	z
  	       |rt!        |j                  |      \  }}nd
\  }}|r |j#                         j	                  |      } nXt%        j&                  |      rdt%        j(                  |	|      z
  nd|	|z  z
  }!||!z  j#                         j	                  |      } || z  }"|r$|"j+                         j	                  |      }#|"|#z  }"|dk(  r)t-        |      dk\  r|rd}$|dd D ]  }%|$|%z  }$	 ||$dz  z  }t/        ||"| |z         / y)u)  Single tensor AdaMuon update.

    AdaMuon applies second-moment estimation to the orthogonalized directions,
    then rescales using RMS-alignment to maintain stable step sizes.

    Algorithm:
        1. Update momentum buffer: M = β₁·M + (1-β₁)·G
        2. Orthogonalize: O = Newton-Schulz(M) or Newton-Schulz(nesterov_update)
        3. Update second moment: v = β₂·v + (1-β₂)·O²
        4. Bias correct: v̂ = v/(1-β₂^t)
        5. Adaptive scaling: Ô = O / (√v̂ + ε)
        6. RMS-aligned rescaling and apply update
    r
   r*   r3   r   r   r   r2   r   )value)r*   FNri   )r   r   r   r   rC   r   r   r   r   r;   r   r!   rY   r   r   addcmul_rj   sqrtrG   	is_tensorpowrA   r<   r   )&r   r   r   r   r   r   r   r   r   r   r   r   r    r.   r[   r   r   r0   r   r   r   r   r   
exp_avg_sqstep_tstepr   r   r   r   r   use_rms_normdenombias_correction2update_adaptiveupdate_normspatial_prodrz   s&                                         r   r   r     s   F .o}MO!&+6f% M:5Qx$Q' ^
Q 	A#vF); 	

1rL(() 	4h/7?L(3\EWEWEY ;;!.>vI.V+O^$O#\\N 2'
 	!l&7&71&<'//1a8L#++N; 	''l#PU+'V "6|7I7I<"XE<",E<OO%**3/E @Et?TsUYYud%;;Z]`eim`mZm"2288:??DE '. )..055c:K-;O 	!c.&9Q&>  '+ &A A%L&-- 	UObS5[9[M:r   c            '            e Zd ZdZddddededdd	d
dddddddfdededededede	de
dededee   dedede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 )%Muona  Muon - MomentUm Orthogonalized by Newton-schulz

    Combines Muon for 2D+ parameters (weight matrices) with AdamW for 1D parameters (biases, norms) and
    parameter groups with 'use_fallback=True' set (or 'use_muon=False' for compatibility).

    Supports two algorithms:
    - "muon": Standard Muon algorithm with momentum + orthogonalization
    - "adamuon": AdaMuon algorithm that adds element-wise second moment estimation
                 to orthogonalized directions for Adam-like adaptive scaling
    g{Gz?r   ffffff?Fr   r*   r^   r   TN)g?r   r   r   r   r   r   r   r   r   r    r.   r[   r   r   adamw_lrfallback_lr_scalebetasalgor0   r   verbosec                    t        d|       t        d|       t        d|d       t        d|       |dvrt        d|       |d	vrt        d
| d      |Qt        j                  dt        d       t        j                  |      rt        d      |dk(  rt        d      ||z  }t        d"i d|d|d|d|d|d|d|d|	d|
d|d|d|d|d|d|d |d!|}t        | %  ||       y)#u0
   Create Muon optimizer.
        Args:
            params: Iterable of parameters or dicts defining parameter groups
            lr: Learning rate (default: 0.02 for Muon parameters)
            weight_decay: Weight decay coefficient
            momentum: Momentum factor for Muon
            nesterov: Whether to use Nesterov momentum
            ns_steps: Number of Newton-Schulz iterations
            ns_coefficients: Coefficients for NS iteration
            eps: Numerical stability epsilon
            safety_factor: Multiplicative safety factor for NS norm
            adjust_lr_fn: LR adjustment function - "original", "match_rms_adamw", or "rms_to_rms".
                For adamuon mode, can set to None to disable (RMS rescaling handles scaling).
            conv_mode: How to handle convolutions - "flatten" or "batched"
            normalize_spatial: Whether to normalize by sqrt(spatial_size) in batched mode
            fallback_lr_scale: Scale factor applied to lr for AdamW fallback parameters.
                The effective fallback LR is lr * fallback_lr_scale. This ensures the LR scheduler
                properly schedules fallback parameters.
            adamw_lr: Deprecated. Use fallback_lr_scale instead.
            betas: Beta coefficients - (beta1, beta2) where beta1 is used for AdamW fallback
                and beta2 is used for both AdamW fallback and AdaMuon second moment
            algo: Algorithm - "muon" for standard Muon, "adamuon" for AdaMuon with
                adaptive second moment estimation (https://arxiv.org/abs/2507.11005)
            scale_eps: If True, scale epsilon by sqrt(din/dout) in Newton-Schulz for μP
                compatibility (https://arxiv.org/abs/2512.05620)
            capturable: Whether this instance is safe to capture in a CUDA graph. Capturable mode supports
                tensor learning rates and requires optimizer state tensors to live on the parameter device.
            verbose: Log parameter routing decisions (Muon vs AdamW)

        Example:
            ```python
            # Simple usage - automatically uses Muon for 2D+ params, AdamW for 1D
            optimizer = Muon(model.parameters(), lr=0.02)

            # Use AdaMuon algorithm for adaptive scaling
            optimizer = Muon(model.parameters(), lr=6e-4, algo="adamuon")

            # Manual control over parameter groups
            optimizer = Muon([
                {'params': weight_matrices, 'lr': 0.02},
                {'params': biases, 'use_fallback': True, 'lr': 3e-4}, # use AdamW if use_fallback=True
            ])
            ```
        zlearning rater   r   r*   )	max_valueepsilon)r   r   zInvalid conv_mode: )r   r   zInvalid algo: z. Must be 'muon' or 'adamuon'Nzpadamw_lr is deprecated, use fallback_lr_scale=adamw_lr/lr instead. adamw_lr will be removed in a future release.r2   )
stacklevelzHadamw_lr is not supported with tensor lr; use fallback_lr_scale instead.r   z8Cannot compute fallback_lr_scale from adamw_lr when lr=0r   r   r   r   r    r.   r[   r   r   r   r   r   r0   r   r   r&   )
r   r   warningswarnFutureWarningrG   r   dictsuper__init__)selfr   r   r   r   r   r   r   r    r.   r[   r   r   r   r   r   r   r0   r   r   defaults	__class__s                        r   r   zMuon.__init__  s   D 	"-6X=C(2229+>??**~dV3PQRRMM@	 r" !kllQw ![\\ (2 

%
 
 	

 
 ,
 
 (
 &
  
 0
 0
 
 
  
  "!
" #
& 	*r   c                    t         |   |       | j                  D ]g  }|j                  dd       |j                  dd       |j                  dd       d|vs>|j	                  d|d         }|d   d	k7  r||d   z  nd
|d<   i y )Nr   r   r0   Fr   r   r   r   r   r*   )r   __setstate__param_groups
setdefaultpop)r   stategroupr   r   s       r   r   zMuon.__setstate__  s    U#&& 	aEVV,[%0\51"%/ 99Zt=GLT{VWGWXd-C]`)*	ar   c                 F   | j                   j                  dd      r9t        | d      r| j                          nt        | d      r| j	                          d}|$t        j                         5   |       }ddd       | j                   j                  dd      }d}d}|ri nd}| j                  D ]  }|j                  dd	      }g }	g }
g }g }g }g }g }g }g }g }|d
   D ]  }|j                  |j                  j                  rt        d      | j                  |   }d|vrd}|j                  dd      r
d|d<   |r9d}n6d|v r|d   |d<   |r(d}n%|rt        |d      \  }}nt        |d      }||d<   |A|?dj                  d |j                  D              }||vrg ||<   ||   j                  |       |d   }|r|	j                  |       |
j                  |j                         |dz  }d|vr(t        j                   |t
        j"                        |d<   |j                  |d          |dk(  s[d|vr:t%        ||d         |d<   t        j                   |t
        j"                        |d<   |j                  |d          |j                  |d          |j                  |       |j                  |j                         |dz  }d|vrbt%        ||d         |d<   t        j                   |t
        j"                        |d<   t        j                   |t
        j"                        |d<   |j                  |d          |j                  |d          |j                  |d           |	r|dk(  rM|d   \  }}t'        |	|
|||f|d   |d   |d   |d   ||d    |d!   |d"   |d#   |d$   |d%   |d&   |d'   |d   d( nBt)        |	|
||d   |d   |d   |d   |d    |d!   |d"   |d#   |d$   |d%   |d&   |d'   |d   )       |sb|d   \  }}|d   |d*   z  }|d   r%t+        |||||d||||d   |d"   dd|d   d+       t-        ||||g |dd||||d   |d"   dd|d   d,        |rt/        |      dkD  rt0        j3                  d-| d.| d/       i }t5        |j7                               D ])  \  }}|D ]  }||vrg ||<   ||   j                  |       ! + g }t5        |j7                               D ]$  \  }} |j                  | d0t/        |               & t0        j3                  d1d2j                  |              t0        j9                  t:        j<                        rt5        |j7                               D ]  \  }} |d3k(  rd4nd5}!t0        j3                  d6| d7|! d8       | dd9 D ]  }"t0        j3                  d:|"         t/        |       d9kD  s\t0        j3                  d;t/        |       d9z
   d<        |S # 1 sw Y   xY w)=z$Performs a single optimization step.r   F'_accelerator_graph_capture_health_check _cuda_graph_capture_health_checkNr   r   r   r   r   z&Muon does not support sparse gradientsuse_muonuse_fallbackuse_fallback_flaguse_muon_flagT)rm   xc              3   2   K   | ]  }t        |        y wN)str)rp   rw   s     r   rr   zMuon.step.<locals>.<genexpr>U  s     ,ESV,Es   r
   momentum_buffer)memory_formatr   r   r   exp_avgr   r   r   r   r   r   r   r    r.   r[   r   r   r0   r   r   r   )
foreachbeta1r   r   r   r    cautionmaximizer   max_lr)r   amsgradr   r   r   r   r    r   r   r   r   zMuon parameter routing: z Muon, z AdamW=z  Breakdown: , ru   r   AdamWz    z -> :
   z      z      ... and z more)r   gethasattrr   r   rG   enable_gradr   r   	is_sparseRuntimeErrorr   r   joinr!   append
zeros_likepreserve_formatr   r   r   r   r   r<   _loggerinfosorteditemsisEnabledForloggingINFO)#r   closurelossr   
muon_countadamw_countrouting_reasonsr   r   muon_params
muon_gradsmuon_momentum_bufsmuon_exp_avg_sqsmuon_state_stepsadamw_paramsadamw_gradsadamw_exp_avgsadamw_exp_avg_sqsadamw_state_stepspr   reasonsuitable	shape_strr   _r   r   fallback_lrreason_groupsreasonsreason_summaryshapesoptimizer_namer!   s#                                      r   r   z	Muon.step  s7    ==\51tFG<<>AB557""$ !y! --##Iu5 
 '"T&& r	E99VV,D KJ!#!!LKN " "8_ H<66>66##&'OPP

1 U*!Fyy7,1j)"%8F#u,,1*,=j)"%4F #/DQVZ/[,Hf'<Qe'TH,4j) '2v7I$'HH,EQWW,E$E	$O;9;OI6'	299&A !,&&q)%%aff-!OJ )5383C3CAUZUjUj3k/0&--e4E.FG y(!.,6q%:M,NE&M272B2B1TYTiTi2jE,/(//l0CD(//f> !''*&&qvv.1$K U*(21eL6I(Jf+0+;+;AUMbMb+ci(.3.>.>qPUPePe.fl+"))%	*:;%,,U<-@A%,,U6];QH<V 9$$W~HAu#"*(( !;%*>%:!&z!2!&z!2#!&z!2(-.?(@!%L&+O&<%*>%:"'"4*/0C*D"'"4#(#6', #"* ;%*>%:!&z!2!&z!2!&z!2(-.?(@!%L&+O&<%*>%:"'"4*/0C*D"'"4#(#6!( $W~u#DkE2E,FF$$#&)) $##&%*>%:!%L %!&#(#6#$ $#&)) $ %##&%*>%:!%L %!&#(#6##Ar	j s?3a7LL3J<w{mSYZ[ M&,_-B-B-D&E <"	7% <F]202f-!&)00;<<  N"()<)<)>"? A%%#f+&?@ALL=>)B(CDE ##GLL1&,]-@-@-B&C ONFF/5~V7NLL4xtN3C1!EF!' 7veW%5676{R'~c&kB6F5Gu%MNO w! !s   .XX r   )__name__
__module____qualname____doc__DEFAULT_NS_STEPSMUON_EPSr   floatboolintNSCoeffr   r   r   r   r   rG   no_gradr   __classcell__)r   s   @r   r   r     sK   	 "#"",'0!#&*;&&*(,'*)4#$!)k+k+ k+  	k+
 k+ k+ k+ %k+ k+ !k+ #3-k+ k+  $k+ uok+  %k+  &!k+" #k+$ %k+& 'k+( )k+Z	a U]]_e er   r   r   presetsc                   	 d 	d dt         t           dt        t        t        t        f   f	fd}t        | t              r}| |vr9dj                  t        |j                                     }t        d|  d|       ||    } 	|      rt        |      d	k(  rt        d
|  d      |D cg c]
  } ||       c}S  	|       st        d      t        |       dk(  rt        fd| D              r	 ||       gS g }t        |       D ]5  \  }} 	|      st        d| d|      |j                   ||             7 |st        d      |S c c}w )Nc                 T    t        | t              xr t        | t        t        f       S r   )rF   r   r   bytesr   s    r   <lambda>z)resolve_ns_coefficients.<locals>.<lambda>  s     z!X.Rz!c5\7R3R r   c                 \    t        | t        j                        xr t        | t               S r   )rF   numbersRealr'  r0  s    r   r1  z)resolve_ns_coefficients.<locals>.<lambda>  s     
1gll3OJq$<O8O r   r   r   c                      |       r"t        |       dk7  st        fd| D              st        d|       | \  }}}t        |      t        |      t        |      fS )Nr3   c              3   .   K   | ]  } |        y wr   r&   rp   vis_reals     r   rr   z<resolve_ns_coefficients.<locals>.as_coeff.<locals>.<genexpr>   s     2I!71:2I   z3Coefficient must be length-3 of real numbers, got: )r<   allr   r&  )r   rR   rS   rT   r9  is_seqs       r   as_coeffz)resolve_ns_coefficients.<locals>.as_coeff  sZ    ayCFaKs2Iq2I/IRSTRWXYY1aQxq58++r   r   zUnknown coefficients preset 'z'. Valid options: r   zPreset 'z' is empty or invalidz]Coefficients must be a preset name (str), a 3-sequence (a,b,c), or a sequence of 3-sequences.r3   c              3   .   K   | ]  } |        y wr   r&   r7  s     r   rr   z*resolve_ns_coefficients.<locals>.<genexpr>  s     9awqz9r:  zItem z is not a sequence: z Coefficient list cannot be empty)r   r&  r   rF   r   r   r  keysr   r<   	TypeErrorr;  r   r   )
r   r,  r=  validseqitemr:   r   r9  r<  s
           @@r   r   r     sj   
 SFOG,HUO ,eUE.A(B , %IIfW\\^45E<UGCUV[U\]^^enc{c#h!mxw.CDEE+./4//%=,
 	
 5zQ39599   CU# #4d|eA3&:4(CDD

8D>"# ;<<J) 0s   0E)F)r^   )   g      `@F)r   )8r#  r  r3  r   typingr   r   r   r   r   r   rG   torch.distributed.tensorr	   rE   ImportError_helpersr   r   r   r   r   _typesr   r   r   	getLoggerr   r   r%  r$  r   r   r&  r)  Tensorr'  r   r(  r)   bfloat16r/   rY   Sizerf   rj   r   r   r   r   r   r   optim	Optimizerr   r   r&   r   r   <module>rP     s  ,    B B 0K k j   
'

H
%  
 	"5&R U5%./eE5%<O6P1QQ
RFell F F F%%S#X% %8 ""^^T<<TT 5u!456T 	T
 T {{T T \\Tr .LZZLL L> .OZZOO 5$;O: "&#	E3||E3E3  E3 	E3
 4tSy!!"E3T  2 2 2 5<<#$ 2h !#%U\\"%ELL!% ELL)%
 % % % % % !% % % sm% %  %  !%" #%$ 
%%x !)2U\\"2ELL!2 ELL)2 %,,'	2
 %,,'2 2 2 2 2 2 2 !2 2 2  sm!2" #2$  %2& '2( )2* 
+2L !#C7U\\"C7ELL!C7 ELL)C7
 C7 C7 C7 C7 C7 !C7 C7 C7 smC7 C7  C7  !C7" #C7$ 
%C7t !)t:U\\"t:ELL!t: ELL)t: %,,'	t:
 %,,'t: t: t: t: t: t: t: !t: t: t:  sm!t:" #t:$  %t:& 't:( )t:* 
+t:nj5;;   jZ)S(5/8HUO+DDE)hx778) 
%ue#
$%)o  Ks   M M)(M)