219219}
220220
221221
222+ # deepseek4 building blocks
223+ _DEEPSEEK4_MHC_ATTENTION = {
224+ "mhc_norm" : {"scale" : None },
225+ "post_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
226+ "post_alpha_scale" : None ,
227+ "post_beta" : None ,
228+ "pre_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
229+ "pre_alpha_scale" : None ,
230+ "pre_beta" : None ,
231+ "res_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
232+ "res_alpha_scale" : None ,
233+ "res_beta" : None ,
234+ }
235+
236+ _DEEPSEEK4_MHC_MLP = {
237+ "mhc_norm" : {"scale" : None },
238+ "post_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
239+ "post_alpha_scale" : None ,
240+ "post_beta" : None ,
241+ "pre_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
242+ "pre_alpha_scale" : None ,
243+ "pre_beta" : None ,
244+ "res_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
245+ "res_alpha_scale" : None ,
246+ "res_beta" : None ,
247+ }
248+
249+ _DEEPSEEK4_MLP = {
250+ "MoeBlock_0" : {
251+ "gate" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
252+ "wi_0" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
253+ "wi_1" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
254+ "wo" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
255+ },
256+ "shared_experts" : {
257+ "wi_0" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
258+ "wi_1" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
259+ "wo" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
260+ },
261+ }
262+
263+ _DEEPSEEK4_MLP_SCANNED = {
264+ "MoeBlock_0" : {
265+ "gate" : {
266+ "bias" : None ,
267+ "kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
268+ },
269+ "wi_0" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
270+ "wi_1" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
271+ "wo" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
272+ },
273+ "shared_experts" : {
274+ "wi_0" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
275+ "wi_1" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
276+ "wo" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
277+ },
278+ }
279+
280+ _DEEPSEEK4_ATTN_BASIC = {
281+ "kv_norm" : {"scale" : None },
282+ "o_a_proj" : {"kernel" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,))},
283+ "o_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
284+ "q_norm" : {"scale" : None },
285+ "sinks" : None ,
286+ "wkv" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 2 , - 1 ))},
287+ "wq_a" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
288+ "wq_b" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 2 , - 1 ))},
289+ }
290+
291+ _DEEPSEEK4_ATTN_CSA = {
292+ "csa_compressor" : {
293+ "gate_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
294+ "indexer" : {
295+ "gate_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
296+ "kv_norm" : {"scale" : None },
297+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
298+ "position_bias" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
299+ "q_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
300+ "weights_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
301+ },
302+ "kv_norm" : {"scale" : None },
303+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
304+ "position_bias" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
305+ },
306+ "kv_norm" : {"scale" : None },
307+ "o_a_proj" : {"kernel" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,))},
308+ "o_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
309+ "q_norm" : {"scale" : None },
310+ "sinks" : None ,
311+ "wkv" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 2 , - 1 ))},
312+ "wq_a" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
313+ "wq_b" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 2 , - 1 ))},
314+ }
315+
316+ _DEEPSEEK4_ATTN_HCA = {
317+ "hca_compressor" : {
318+ "gate_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
319+ "kv_norm" : {"scale" : None },
320+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
321+ "position_bias" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
322+ },
323+ "kv_norm" : {"scale" : None },
324+ "o_a_proj" : {"kernel" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,))},
325+ "o_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
326+ "q_norm" : {"scale" : None },
327+ "sinks" : None ,
328+ "wkv" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 2 , - 1 ))},
329+ "wq_a" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
330+ "wq_b" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 2 , - 1 ))},
331+ }
332+
333+ _DEEPSEEK4_LAYER_BASIC = {
334+ "mhc_attention" : _DEEPSEEK4_MHC_ATTENTION ,
335+ "mhc_mlp" : _DEEPSEEK4_MHC_MLP ,
336+ "mlp" : _DEEPSEEK4_MLP ,
337+ "post_self_attention_layer_norm" : {"scale" : None },
338+ "pre_self_attention_layer_norm" : {"scale" : None },
339+ "self_attention" : _DEEPSEEK4_ATTN_BASIC ,
340+ }
341+
342+ _DEEPSEEK4_LAYER_CSA_PREFIX = {
343+ "mhc_attention" : _DEEPSEEK4_MHC_ATTENTION ,
344+ "mhc_mlp" : _DEEPSEEK4_MHC_MLP ,
345+ "mlp" : _DEEPSEEK4_MLP ,
346+ "post_self_attention_layer_norm" : {"scale" : None },
347+ "pre_self_attention_layer_norm" : {"scale" : None },
348+ "self_attention" : _DEEPSEEK4_ATTN_CSA ,
349+ }
350+
351+ _DEEPSEEK4_LAYER_CSA_SCANNED = {
352+ "mhc_attention" : _DEEPSEEK4_MHC_ATTENTION ,
353+ "mhc_mlp" : _DEEPSEEK4_MHC_MLP ,
354+ "mlp" : _DEEPSEEK4_MLP_SCANNED ,
355+ "post_self_attention_layer_norm" : {"scale" : None },
356+ "pre_self_attention_layer_norm" : {"scale" : None },
357+ "self_attention" : _DEEPSEEK4_ATTN_CSA ,
358+ }
359+
360+ _DEEPSEEK4_LAYER_HCA_SCANNED = {
361+ "mhc_attention" : _DEEPSEEK4_MHC_ATTENTION ,
362+ "mhc_mlp" : _DEEPSEEK4_MHC_MLP ,
363+ "mlp" : _DEEPSEEK4_MLP_SCANNED ,
364+ "post_self_attention_layer_norm" : {"scale" : None },
365+ "pre_self_attention_layer_norm" : {"scale" : None },
366+ "self_attention" : _DEEPSEEK4_ATTN_HCA ,
367+ }
368+
369+ DEEPSEEK4_DIMENSION_NUMBER = {
370+ "params" : {
371+ "decoder" : {
372+ "decoder_norm" : {"scale" : None },
373+ "hc_head" : {
374+ "hc_base" : None ,
375+ "hc_fn" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
376+ "hc_scale" : None ,
377+ },
378+ "layers_0" : _DEEPSEEK4_LAYER_BASIC ,
379+ "layers_1" : _DEEPSEEK4_LAYER_BASIC ,
380+ "layers_2" : _DEEPSEEK4_LAYER_CSA_PREFIX ,
381+ "logits_dense" : {"kernel" : None },
382+ "scanned_blocks" : {
383+ "layers_0" : _DEEPSEEK4_LAYER_HCA_SCANNED ,
384+ "layers_1" : _DEEPSEEK4_LAYER_CSA_SCANNED ,
385+ },
386+ },
387+ "token_embedder" : {"embedding" : None },
388+ }
389+ }
390+
391+
222392class MuonDimensionTest (parameterized .TestCase ):
223393 """Unit tests for Muon dimension number generation.
224394
@@ -229,6 +399,7 @@ class MuonDimensionTest(parameterized.TestCase):
229399 @parameterized .named_parameters (
230400 ("deepseek2-16b" , "deepseek2-16b" , DEEPSEEK2_DIMENSION_NUMBER ),
231401 ("deepseek3-671b" , "deepseek3-671b" , DEEPSEEK3_DIMENSION_NUMBER ),
402+ ("deepseek4-284b" , "deepseek4-284b" , DEEPSEEK4_DIMENSION_NUMBER ),
232403 ("kimi-k2-1t" , "kimi-k2-1t" , DEEPSEEK3_DIMENSION_NUMBER ),
233404 ("llama2-7b" , "llama2-7b" , LLAMA2_DIMENSION_NUMBER ),
234405 ("llama3-8b" , "llama3-8b" , LLAMA2_DIMENSION_NUMBER ),
@@ -244,7 +415,10 @@ def test_model_integration(self, model_name, expected_output):
244415 Muon dimension numbers match the hardcoded reference.
245416 """
246417 actual_output = muon_utils .get_model_mdn (model_name , scan_layers = True , pure_nnx = False )
247- self .assertEqual (actual_output , expected_output )
418+ if "params" in expected_output and "params" in actual_output :
419+ self .assertEqual (actual_output ["params" ], expected_output ["params" ])
420+ else :
421+ self .assertEqual (actual_output , expected_output )
248422
249423
250424class AdamWMaskTest (parameterized .TestCase ):
@@ -621,6 +795,49 @@ def __init__(self, rngs: nnx.Rngs):
621795 # Check attention out: [0, -2] -> [-1]
622796 self .assertEqual (result .self_attention .out .kernel , mdn ((0 , - 2 ), (- 1 ,)))
623797
798+ def test_muon_ds4_ns_config (self ):
799+ """Verifies that muon optimizer configures Newton-Schulz parameters correctly based on model."""
800+ model = MagicMock ()
801+ learning_rate_schedule = MagicMock ()
802+
803+ # Case 1: DeepSeek4 Model (Auto-configures 10-step schedule)
804+ argv_ds4 = [
805+ "" ,
806+ get_test_config_path (),
807+ "run_name=test" ,
808+ "opt_type=muon" ,
809+ "model_name=deepseek4-284b" ,
810+ "attention=dot_product" ,
811+ ]
812+ config_ds4 = pyconfig .initialize (argv_ds4 )
813+
814+ with (
815+ patch ("maxtext.optimizers.optimizers.get_muon_weight_dimension_numbers" ) as mock_get_mdn ,
816+ patch ("maxtext.optimizers.optimizers.muon" ) as mock_muon ,
817+ ):
818+ mock_get_mdn .return_value = {}
819+ optimizers .get_optimizer (config_ds4 , learning_rate_schedule , model = model )
820+ mock_muon .assert_called_once ()
821+ _ , kwargs = mock_muon .call_args
822+ self .assertEqual (kwargs ["ns_steps" ], 10 )
823+ self .assertEqual (len (kwargs ["ns_coeffs" ]), 10 )
824+ self .assertEqual (kwargs ["ns_coeffs" ][- 1 ], (2.0 , - 1.5 , 0.5 ))
825+
826+ # Case 2: Standard Model (Llama2) (Defaults to 5-step schedule)
827+ argv_llama = ["" , get_test_config_path (), "run_name=test" , "opt_type=muon" , "model_name=llama2-7b" ]
828+ config_llama = pyconfig .initialize (argv_llama )
829+
830+ with (
831+ patch ("maxtext.optimizers.optimizers.get_muon_weight_dimension_numbers" ) as mock_get_mdn ,
832+ patch ("maxtext.optimizers.optimizers.muon" ) as mock_muon ,
833+ ):
834+ mock_get_mdn .return_value = {}
835+ optimizers .get_optimizer (config_llama , learning_rate_schedule , model = model )
836+ mock_muon .assert_called_once ()
837+ _ , kwargs = mock_muon .call_args
838+ self .assertEqual (kwargs ["ns_steps" ], 5 )
839+ self .assertEqual (kwargs ["ns_coeffs" ], (3.4445 , - 4.7750 , 2.0315 ))
840+
624841
625842if __name__ == "__main__" :
626843 unittest .main ()
0 commit comments