@@ -39,25 +39,25 @@ def zeropower_via_newtonschulz5(G: Tensor, steps: int, use_bf16: bool) -> Tensor
3939 """
4040 assert G .ndim == 3 # batched Muon implementation by @scottjmaddox, and put into practice in the record by @YouJiacheng
4141 a , b , c = (3.4445 , - 4.7750 , 2.0315 )
42- if use_bf16 :
43- X = G .bfloat16 ()
44- else :
45- X = G .float ()
46- if G .size (- 2 ) > G .size (- 1 ):
47- X = X .mT
42+
43+ X = G .to (dtype = torch .bfloat16 if use_bf16 else torch .float32 )
4844
4945 # Ensure spectral norm is at most 1
5046 X = F .normalize (X , p = 2.0 , dim = (- 2 , - 1 ), eps = 1e-7 )
5147
5248 # Perform the NS iterations
53- for _ in range (steps ):
54- A = X @ X .mT
55- B = torch .baddbmm (A , A , A , beta = b , alpha = c )
56- X = torch .baddbmm (X , B , X , beta = a , alpha = 1 )
49+ if X .size (- 2 ) < X .size (- 1 ):
50+ for _ in range (steps ):
51+ A = torch .bmm (X , X .mT )
52+ A = torch .baddbmm (A , A , A , beta = b , alpha = c )
53+ X = torch .baddbmm (X , A , X , beta = a , alpha = 1 )
54+ else :
55+ for _ in range (steps ):
56+ A = torch .bmm (X .mT , X )
57+ A = torch .baddbmm (A , A , A , beta = b , alpha = c )
58+ X = torch .baddbmm (X , X , A , beta = a , alpha = 1 )
5759
58- if G .size (- 2 ) > G .size (- 1 ):
59- X = X .mT
60- return X .to (G )
60+ return X
6161
6262
6363class Muon (torch .optim .Optimizer ):
@@ -92,33 +92,34 @@ def __init__(self, params, lr=5e-4, weight_decay=0.1, momentum=0.95, nesterov=Tr
9292 def step (self , closure = None ):
9393 for group in self .param_groups :
9494 shape_groups = {}
95- for p in filter (lambda p : p .grad is not None , group ["params" ]):
95+ for p in filter (lambda _p : _p .grad is not None , group ["params" ]):
9696 g = p .grad
9797 state = self .state [p ]
9898 if "momentum_buffer" not in state :
9999 state ["momentum_buffer" ] = torch .zeros_like (g )
100- buf : Tensor = state ["momentum_buffer" ]
101100 key = (p .shape , p .device , p .dtype )
102101 if key not in shape_groups :
103102 shape_groups [key ] = {"params" : [], "grads" : [], "buffers" : []}
104103 shape_groups [key ]["params" ].append (p )
105104 shape_groups [key ]["grads" ].append (g )
106- shape_groups [key ]["buffers" ].append (buf )
105+ shape_groups [key ]["buffers" ].append (state [ "momentum_buffer" ] )
107106 for key in shape_groups :
108107 group_data = shape_groups [key ]
109- g = torch .stack (group_data ["grads" ])
110- buf = torch .stack (group_data ["buffers" ])
111- buf .lerp_ (g , 1 - group ["momentum" ])
112- g = g .lerp_ (buf , group ["momentum" ]) if group ["nesterov" ] else buf
108+ p , g , buf , m = group_data ["params" ], group_data ["grads" ], group_data ["buffers" ], group ["momentum" ]
109+ torch ._foreach_lerp_ (buf , g , 1 - m )
110+ if group ["nesterov" ]:
111+ torch ._foreach_lerp_ (g , buf , m )
112+ g = torch .stack (g )
113+ else :
114+ g = torch .stack (buf )
115+ original_shape = g .shape
113116 if g .ndim >= 4 : # for the case of conv filters
114117 g = g .view (g .size (0 ), g .size (1 ), - 1 )
115118 use_bf16 = self .bf16_support_map .get (g .device , False )
116119 g = zeropower_via_newtonschulz5 (g , steps = group ["ns_steps" ], use_bf16 = use_bf16 )
117- for i , p in enumerate (group_data ["params" ]):
118- if group ["weight_decay" ] > 0 :
119- p .data .mul_ (1 - group ["lr" ] * group ["weight_decay" ])
120- p .data .add_ (g [i ].view_as (p ), alpha = - group ["lr" ] * max (g [i ].size ()) ** 0.5 )
121- self .state [p ]["momentum_buffer" ] = buf [i ].clone ()
120+ if group ["weight_decay" ] > 0 :
121+ torch ._foreach_mul_ (p , 1 - group ["lr" ] * group ["weight_decay" ])
122+ torch ._foreach_add_ (p , g .view (original_shape ).unbind (0 ), alpha = - group ["lr" ] * max (g [0 ].size ()) ** 0.5 )
122123
123124
124125def get_params_for_muon (model ) -> List [Parameter ]:
@@ -129,6 +130,7 @@ def get_params_for_muon(model) -> List[Parameter]:
129130 module: The module to filter parameters for.
130131 Returns:
131132 A list of parameters that should be optimized with muon.
133+ :param model:
132134 """
133135 muon_params = []
134136 for module in model .modules ():
@@ -141,7 +143,11 @@ def get_params_for_muon(model) -> List[Parameter]:
141143
142144
143145class Muon_AdamW (ChainedOptimizer ):
144- def __init__ (self , model , lr = 0.0005 , weight_decay = 0.0 , muon_args = {}, adamw_args = {}, verbose = False ):
146+ def __init__ (self , model , lr = 0.0005 , weight_decay = 0.0 , muon_args = None , adamw_args = None , verbose = False ):
147+ if adamw_args is None :
148+ adamw_args = {}
149+ if muon_args is None :
150+ muon_args = {}
145151 muon_params_id_set = set (id (p ) for p in get_params_for_muon (model ))
146152 spec_muon = OptimizerSpec (Muon , muon_args , lambda param : id (param ) in muon_params_id_set )
147153 spec_adamw = OptimizerSpec (torch .optim .AdamW , adamw_args , None )
0 commit comments