219219}
220220
221221
222+ # deepseek_v4 building blocks
223+ _DEEPSEEK_V4_MHC_ATTENTION = {
224+ "post_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
225+ "post_alpha_scale" : None ,
226+ "post_beta" : None ,
227+ "pre_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
228+ "pre_alpha_scale" : None ,
229+ "pre_beta" : None ,
230+ "res_alpha" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
231+ "res_alpha_scale" : None ,
232+ "res_beta" : None ,
233+ }
234+
235+ _DEEPSEEK_V4_MHC_MLP = _DEEPSEEK_V4_MHC_ATTENTION
236+
237+ _DEEPSEEK_V4_MLP_STD = {
238+ "MoeBlock_0" : {
239+ "gate" : {
240+ "e_score_correction_bias" : None ,
241+ "kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
242+ },
243+ "wi_0" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
244+ "wi_1" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
245+ "wo" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
246+ },
247+ "shared_experts" : {
248+ "wi" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
249+ "wo" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
250+ },
251+ }
252+
253+ _DEEPSEEK_V4_MLP_PRE = {
254+ "MoeBlock_0" : {
255+ "gate" : {
256+ "kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
257+ "tid2eid" : None ,
258+ },
259+ "wi_0" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
260+ "wi_1" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
261+ "wo" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,)),
262+ },
263+ "shared_experts" : {
264+ "wi" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
265+ "wo" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
266+ },
267+ }
268+
269+ _DEEPSEEK_V4_ATTN_INDEXER = {
270+ "compressor" : {
271+ "gate_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
272+ "indexer" : {
273+ "gate_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
274+ "kv_norm" : {"scale" : None },
275+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
276+ "position_bias" : None ,
277+ "q_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
278+ "weights_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
279+ },
280+ "kv_norm" : {"scale" : None },
281+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
282+ "position_bias" : None ,
283+ },
284+ "kv_norm" : {"scale" : None },
285+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
286+ "o_a_proj" : {"kernel" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,))},
287+ "o_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
288+ "q_a_norm" : {"scale" : None },
289+ "q_a_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
290+ "q_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
291+ "sinks" : None ,
292+ }
293+
294+ _DEEPSEEK_V4_ATTN_COMPRESSOR = {
295+ "compressor" : {
296+ "gate_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
297+ "kv_norm" : {"scale" : None },
298+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
299+ "position_bias" : None ,
300+ },
301+ "kv_norm" : {"scale" : None },
302+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
303+ "o_a_proj" : {"kernel" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,))},
304+ "o_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
305+ "q_a_norm" : {"scale" : None },
306+ "q_a_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
307+ "q_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
308+ "sinks" : None ,
309+ }
310+
311+ _DEEPSEEK_V4_ATTN_BASIC = {
312+ "kv_norm" : {"scale" : None },
313+ "kv_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
314+ "o_a_proj" : {"kernel" : mdn (reduction_axis = (- 2 ,), output_axis = (- 1 ,))},
315+ "o_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
316+ "q_a_norm" : {"scale" : None },
317+ "q_a_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
318+ "q_b_proj" : {"kernel" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,))},
319+ "sinks" : None ,
320+ }
321+
322+ _DEEPSEEK_V4_LAYER_INDEXER = {
323+ "mhc_attention" : _DEEPSEEK_V4_MHC_ATTENTION ,
324+ "mhc_mlp" : _DEEPSEEK_V4_MHC_MLP ,
325+ "mlp" : _DEEPSEEK_V4_MLP_STD ,
326+ "post_self_attention_layer_norm" : {"scale" : None },
327+ "pre_self_attention_layer_norm" : {"scale" : None },
328+ "self_attention" : _DEEPSEEK_V4_ATTN_INDEXER ,
329+ }
330+
331+ _DEEPSEEK_V4_LAYER_COMPRESSOR = {
332+ "mhc_attention" : _DEEPSEEK_V4_MHC_ATTENTION ,
333+ "mhc_mlp" : _DEEPSEEK_V4_MHC_MLP ,
334+ "mlp" : _DEEPSEEK_V4_MLP_STD ,
335+ "post_self_attention_layer_norm" : {"scale" : None },
336+ "pre_self_attention_layer_norm" : {"scale" : None },
337+ "self_attention" : _DEEPSEEK_V4_ATTN_COMPRESSOR ,
338+ }
339+
340+ _DEEPSEEK_V4_LAYER_BASIC_STD = {
341+ "mhc_attention" : _DEEPSEEK_V4_MHC_ATTENTION ,
342+ "mhc_mlp" : _DEEPSEEK_V4_MHC_MLP ,
343+ "mlp" : _DEEPSEEK_V4_MLP_STD ,
344+ "post_self_attention_layer_norm" : {"scale" : None },
345+ "pre_self_attention_layer_norm" : {"scale" : None },
346+ "self_attention" : _DEEPSEEK_V4_ATTN_BASIC ,
347+ }
348+
349+ _DEEPSEEK_V4_LAYER_BASIC_PRE = {
350+ "mhc_attention" : _DEEPSEEK_V4_MHC_ATTENTION ,
351+ "mhc_mlp" : _DEEPSEEK_V4_MHC_MLP ,
352+ "mlp" : _DEEPSEEK_V4_MLP_PRE ,
353+ "post_self_attention_layer_norm" : {"scale" : None },
354+ "pre_self_attention_layer_norm" : {"scale" : None },
355+ "self_attention" : _DEEPSEEK_V4_ATTN_BASIC ,
356+ }
357+
358+ _DEEPSEEK_V4_LAYER_INDEXER_PRE = {
359+ "mhc_attention" : _DEEPSEEK_V4_MHC_ATTENTION ,
360+ "mhc_mlp" : _DEEPSEEK_V4_MHC_MLP ,
361+ "mlp" : _DEEPSEEK_V4_MLP_PRE ,
362+ "post_self_attention_layer_norm" : {"scale" : None },
363+ "pre_self_attention_layer_norm" : {"scale" : None },
364+ "self_attention" : _DEEPSEEK_V4_ATTN_INDEXER ,
365+ }
366+
367+ _DEEPSEEK_V4_LAYER_COMPRESSOR_PRE = {
368+ "mhc_attention" : _DEEPSEEK_V4_MHC_ATTENTION ,
369+ "mhc_mlp" : _DEEPSEEK_V4_MHC_MLP ,
370+ "mlp" : _DEEPSEEK_V4_MLP_PRE ,
371+ "post_self_attention_layer_norm" : {"scale" : None },
372+ "pre_self_attention_layer_norm" : {"scale" : None },
373+ "self_attention" : _DEEPSEEK_V4_ATTN_COMPRESSOR ,
374+ }
375+
376+ DEEPSEEK_V4_DIMENSION_NUMBER = {
377+ "params" : {
378+ "decoder" : {
379+ "decoder_norm" : {"scale" : None },
380+ "hc_head" : {
381+ "hc_base" : None ,
382+ "hc_fn" : mdn (reduction_axis = (0 ,), output_axis = (- 1 ,)),
383+ "hc_scale" : None ,
384+ },
385+ "layers" : {
386+ "layers_0" : _DEEPSEEK_V4_LAYER_INDEXER ,
387+ "layers_1" : _DEEPSEEK_V4_LAYER_COMPRESSOR ,
388+ },
389+ "logits_dense" : {"kernel" : None },
390+ "post_layers" : {
391+ "layers_0" : _DEEPSEEK_V4_LAYER_BASIC_STD ,
392+ },
393+ "pre_layers" : {
394+ "layers_0" : _DEEPSEEK_V4_LAYER_BASIC_PRE ,
395+ "layers_1" : _DEEPSEEK_V4_LAYER_INDEXER_PRE ,
396+ "layers_2" : _DEEPSEEK_V4_LAYER_COMPRESSOR_PRE ,
397+ },
398+ },
399+ "token_embedder" : {"embedding" : None },
400+ }
401+ }
402+
403+
222404class MuonDimensionTest (parameterized .TestCase ):
223405 """Unit tests for Muon dimension number generation.
224406
@@ -229,6 +411,7 @@ class MuonDimensionTest(parameterized.TestCase):
229411 @parameterized .named_parameters (
230412 ("deepseek2-16b" , "deepseek2-16b" , DEEPSEEK2_DIMENSION_NUMBER ),
231413 ("deepseek3-671b" , "deepseek3-671b" , DEEPSEEK3_DIMENSION_NUMBER ),
414+ ("deepseek_v4-tiny" , "deepseek_v4-tiny" , DEEPSEEK_V4_DIMENSION_NUMBER ),
232415 ("kimi-k2-1t" , "kimi-k2-1t" , DEEPSEEK3_DIMENSION_NUMBER ),
233416 ("llama2-7b" , "llama2-7b" , LLAMA2_DIMENSION_NUMBER ),
234417 ("llama3-8b" , "llama3-8b" , LLAMA2_DIMENSION_NUMBER ),
@@ -244,7 +427,10 @@ def test_model_integration(self, model_name, expected_output):
244427 Muon dimension numbers match the hardcoded reference.
245428 """
246429 actual_output = muon_utils .get_model_mdn (model_name , scan_layers = True , pure_nnx = False )
247- self .assertEqual (actual_output , expected_output )
430+ if "params" in expected_output and "params" in actual_output :
431+ self .assertEqual (actual_output ["params" ], expected_output ["params" ])
432+ else :
433+ self .assertEqual (actual_output , expected_output )
248434
249435
250436class AdamWMaskTest (parameterized .TestCase ):
@@ -621,6 +807,38 @@ def __init__(self, rngs: nnx.Rngs):
621807 # Check attention out: [0, -2] -> [-1]
622808 self .assertEqual (result .self_attention .out .kernel .value , mdn ((0 , - 2 ), (- 1 ,)))
623809
810+ def test_muon_ds4_ns_config (self ):
811+ """Verifies that muon optimizer configures Newton-Schulz parameters correctly based on ds4_ns."""
812+ model = MagicMock ()
813+ learning_rate_schedule = MagicMock ()
814+
815+ # Case 1: ds4_ns = True
816+ argv_true = ["" , get_test_config_path (), "run_name=test" , "opt_type=muon" , "ds4_ns=true" ]
817+ config_true = pyconfig .initialize (argv_true )
818+
819+ with patch ("maxtext.optimizers.optimizers.get_muon_weight_dimension_numbers" ) as mock_get_mdn , \
820+ patch ("maxtext.optimizers.optimizers.muon" ) as mock_muon :
821+ mock_get_mdn .return_value = {}
822+ optimizers .get_optimizer (config_true , learning_rate_schedule , model = model )
823+ mock_muon .assert_called_once ()
824+ _ , kwargs = mock_muon .call_args
825+ self .assertEqual (kwargs ["ns_steps" ], 10 )
826+ self .assertEqual (len (kwargs ["ns_coeffs" ]), 10 )
827+ self .assertEqual (kwargs ["ns_coeffs" ][- 1 ], (2.0 , - 1.5 , 0.5 ))
828+
829+ # Case 2: ds4_ns = False (Default)
830+ argv_false = ["" , get_test_config_path (), "run_name=test" , "opt_type=muon" , "ds4_ns=false" ]
831+ config_false = pyconfig .initialize (argv_false )
832+
833+ with patch ("maxtext.optimizers.optimizers.get_muon_weight_dimension_numbers" ) as mock_get_mdn , \
834+ patch ("maxtext.optimizers.optimizers.muon" ) as mock_muon :
835+ mock_get_mdn .return_value = {}
836+ optimizers .get_optimizer (config_false , learning_rate_schedule , model = model )
837+ mock_muon .assert_called_once ()
838+ _ , kwargs = mock_muon .call_args
839+ self .assertEqual (kwargs ["ns_steps" ], 5 )
840+ self .assertEqual (kwargs ["ns_coeffs" ], (3.4445 , - 4.7750 , 2.0315 ))
841+
624842
625843if __name__ == "__main__" :
626844 unittest .main ()
0 commit comments