From fecc14bbccb64cc2c8092e213756c0709de3c63f Mon Sep 17 00:00:00 2001 From: Vadim Smirnov Date: Sat, 4 Jul 2026 11:53:07 +0400 Subject: [PATCH 1/3] fix: preserve qwen3.5 moe switch projection names --- src/models/qwen3_5.rs | 25 ------------------------- src/models/qwen3_5_tests.rs | 32 +++++++++++++++++++++++++++++++- 2 files changed, 31 insertions(+), 26 deletions(-) diff --git a/src/models/qwen3_5.rs b/src/models/qwen3_5.rs index bd5a3efa1..3d4b1fc59 100644 --- a/src/models/qwen3_5.rs +++ b/src/models/qwen3_5.rs @@ -2481,31 +2481,6 @@ pub fn sanitize_weights(mut weights: WeightMap, config: &Qwen35Config) -> Weight } } - // 8. Rename switch_mlp.{gate_proj,up_proj,down_proj} → switch_mlp.{w1,w3,w2} - // Pre-quantized MoE models use gate_proj/up_proj/down_proj naming, - // but SparseMoeBlock expects w1/w2/w3 naming. - let rename_map = [ - ("switch_mlp.gate_proj.", "switch_mlp.w1."), - ("switch_mlp.up_proj.", "switch_mlp.w3."), - ("switch_mlp.down_proj.", "switch_mlp.w2."), - ]; - let keys_to_rename: Vec = weights - .keys() - .filter(|k| rename_map.iter().any(|(from, _)| k.contains(from))) - .cloned() - .collect(); - for key in keys_to_rename { - for (from, to) in &rename_map { - if key.contains(from) { - let new_key = key.replace(from, to); - if let Some(v) = weights.remove(&key) { - weights.insert(new_key, v); - } - break; - } - } - } - weights } diff --git a/src/models/qwen3_5_tests.rs b/src/models/qwen3_5_tests.rs index e61e9af49..d214424a0 100644 --- a/src/models/qwen3_5_tests.rs +++ b/src/models/qwen3_5_tests.rs @@ -21,7 +21,8 @@ //! that need a real Qwen 3.5 model and are gated behind hardware availability. use super::qwen3_5::{ - Qwen35Config, rebuild_with_zero_tail, sanitize_weights, zero_per_row_kv_tail, + Qwen35Config, rebuild_with_zero_tail, sanitize_moe_weights, sanitize_weights, + zero_per_row_kv_tail, }; use mlxcel_core::dtype; use mlxcel_core::layers::KVCache; @@ -309,3 +310,32 @@ fn sanitize_weights_drops_lm_head_when_tied_embeddings() { ); assert!(sanitized.contains_key("model.embed_tokens.weight")); } + +#[test] +#[ignore = "requires serial MLX execution"] +fn sanitize_moe_weights_preserves_stacked_switch_proj_names() { + let root = "language_model.model.layers.0.mlp.switch_mlp"; + let mut weights = WeightMap::new(); + for proj in ["gate_proj", "up_proj", "down_proj"] { + weights.insert( + format!("{root}.{proj}.weight"), + mlxcel_core::from_slice_f32(&[0.0_f32; 8], &[2, 4]), + ); + } + + let mut config = make_tiny_config(); + config.num_experts = 2; + config.num_experts_per_tok = 1; + config.decoder_sparse_step = 1; + config.moe_intermediate_size = 2; + config.shared_expert_intermediate_size = 2; + + let sanitized = sanitize_moe_weights(weights, &config); + + assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.gate_proj.weight")); + assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.up_proj.weight")); + assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.down_proj.weight")); + assert!(!sanitized.contains_key("model.layers.0.mlp.switch_mlp.w1.weight")); + assert!(!sanitized.contains_key("model.layers.0.mlp.switch_mlp.w2.weight")); + assert!(!sanitized.contains_key("model.layers.0.mlp.switch_mlp.w3.weight")); +} From 4a51d20fb7bb55e1cfb928bffb5eaf252f258cf0 Mon Sep 17 00:00:00 2001 From: Vadim Smirnov Date: Sat, 4 Jul 2026 12:21:00 +0400 Subject: [PATCH 2/3] fix: stack qwen3.5 moe experts under loader names --- src/models/qwen3_5.rs | 21 ++++++++++----------- src/models/qwen3_5_tests.rs | 20 +++++++++++--------- 2 files changed, 21 insertions(+), 20 deletions(-) diff --git a/src/models/qwen3_5.rs b/src/models/qwen3_5.rs index 3d4b1fc59..4508c6e27 100644 --- a/src/models/qwen3_5.rs +++ b/src/models/qwen3_5.rs @@ -2398,11 +2398,10 @@ pub fn sanitize_weights(mut weights: WeightMap, config: &Qwen35Config) -> Weight // Handle per-expert gate_proj/up_proj/down_proj naming variant. // Checkpoints that store experts under `experts.{e}.gate_proj.weight` - // instead of `experts.{e}.w1.weight` use this layout. The mapping - // mirrors the fused gate_up_proj path: gate->w1, up->w3, down->w2. - for (src_proj, dst_proj) in [("gate_proj", "w1"), ("up_proj", "w3"), ("down_proj", "w2")] { - // Skip if this target slot is already populated by the w1/w2/w3 pass above. - if weights.contains_key(format!("{}.{}.weight", base, dst_proj).as_str()) { + // instead of per-expert w1 weights use this layout. Keep the + // stacked names aligned with SwitchGLU::from_weights. + for proj in ["gate_proj", "up_proj", "down_proj"] { + if weights.contains_key(format!("{}.{}.weight", base, proj).as_str()) { continue; } @@ -2413,18 +2412,18 @@ pub fn sanitize_weights(mut weights: WeightMap, config: &Qwen35Config) -> Weight let mut e = 0; while let Some(w) = weights.remove(&format!( "model.layers.{}.mlp.experts.{}.{}.weight", - l, e, src_proj + l, e, proj )) { expert_weights.push(w); if let Some(s) = weights.remove(&format!( "model.layers.{}.mlp.experts.{}.{}.scales", - l, e, src_proj + l, e, proj )) { expert_scales.push(s); } if let Some(b) = weights.remove(&format!( "model.layers.{}.mlp.experts.{}.{}.biases", - l, e, src_proj + l, e, proj )) { expert_biases.push(b); } @@ -2433,16 +2432,16 @@ pub fn sanitize_weights(mut weights: WeightMap, config: &Qwen35Config) -> Weight if !expert_weights.is_empty() { let stacked = stack_arrays(&expert_weights, 0); - weights.insert(format!("{}.{}.weight", base, dst_proj), stacked); + weights.insert(format!("{}.{}.weight", base, proj), stacked); if !expert_scales.is_empty() { let stacked = stack_arrays(&expert_scales, 0); - weights.insert(format!("{}.{}.scales", base, dst_proj), stacked); + weights.insert(format!("{}.{}.scales", base, proj), stacked); } if !expert_biases.is_empty() { let stacked = stack_arrays(&expert_biases, 0); - weights.insert(format!("{}.{}.biases", base, dst_proj), stacked); + weights.insert(format!("{}.{}.biases", base, proj), stacked); } } } diff --git a/src/models/qwen3_5_tests.rs b/src/models/qwen3_5_tests.rs index d214424a0..f84a176bd 100644 --- a/src/models/qwen3_5_tests.rs +++ b/src/models/qwen3_5_tests.rs @@ -21,7 +21,7 @@ //! that need a real Qwen 3.5 model and are gated behind hardware availability. use super::qwen3_5::{ - Qwen35Config, rebuild_with_zero_tail, sanitize_moe_weights, sanitize_weights, + Qwen35Config, rebuild_with_zero_tail, sanitize_weights, zero_per_row_kv_tail, }; use mlxcel_core::dtype; @@ -313,14 +313,16 @@ fn sanitize_weights_drops_lm_head_when_tied_embeddings() { #[test] #[ignore = "requires serial MLX execution"] -fn sanitize_moe_weights_preserves_stacked_switch_proj_names() { - let root = "language_model.model.layers.0.mlp.switch_mlp"; +fn sanitize_weights_stacks_per_expert_switch_proj_names_for_loader() { + let root = "model.layers.0.mlp.experts"; let mut weights = WeightMap::new(); - for proj in ["gate_proj", "up_proj", "down_proj"] { - weights.insert( - format!("{root}.{proj}.weight"), - mlxcel_core::from_slice_f32(&[0.0_f32; 8], &[2, 4]), - ); + for expert in 0..2 { + for proj in ["gate_proj", "up_proj", "down_proj"] { + weights.insert( + format!("{root}.{expert}.{proj}.weight"), + mlxcel_core::from_slice_f32(&[expert as f32; 8], &[2, 4]), + ); + } } let mut config = make_tiny_config(); @@ -330,7 +332,7 @@ fn sanitize_moe_weights_preserves_stacked_switch_proj_names() { config.moe_intermediate_size = 2; config.shared_expert_intermediate_size = 2; - let sanitized = sanitize_moe_weights(weights, &config); + let sanitized = sanitize_weights(weights, &config); assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.gate_proj.weight")); assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.up_proj.weight")); From 0cf5d7dd44f83729bccd59de216f03cb3da72eae Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Wed, 8 Jul 2026 08:56:48 +0900 Subject: [PATCH 3/3] test(qwen3.5): preserve MoE expert stack ordering --- src/models/qwen3_5.rs | 62 +++++++++++++++++++++++++++++++++++++ src/models/qwen3_5_tests.rs | 34 +------------------- 2 files changed, 63 insertions(+), 33 deletions(-) diff --git a/src/models/qwen3_5.rs b/src/models/qwen3_5.rs index b760a483b..930630da5 100644 --- a/src/models/qwen3_5.rs +++ b/src/models/qwen3_5.rs @@ -2915,6 +2915,68 @@ mod sanitize_tests { } } + #[test] + fn sanitize_weights_preserves_per_expert_gate_up_down_stack_order() { + let num_experts: usize = 2; + let out = 2i32; + let in_dim = 4i32; + let config = moe_config(1, num_experts); + + let mut weights = WeightMap::new(); + for expert in 0..num_experts { + for (proj, base) in [ + ("gate_proj", 10.0_f32), + ("up_proj", 20.0_f32), + ("down_proj", 30.0_f32), + ] { + let value = base + expert as f32; + weights.insert( + format!("model.layers.0.mlp.experts.{}.{}.weight", expert, proj), + mlxcel_core::from_slice_f32( + &vec![value; (out * in_dim) as usize], + &[out, in_dim], + ), + ); + } + } + + let result = sanitize_weights(weights, &config); + + for (proj, base) in [ + ("gate_proj", 10.0_f32), + ("up_proj", 20.0_f32), + ("down_proj", 30.0_f32), + ] { + let key = format!("model.layers.0.mlp.switch_mlp.{}.weight", proj); + let arr = result + .get(key.as_str()) + .unwrap_or_else(|| panic!("missing stacked weight at {key}")); + assert_eq!( + mlxcel_core::array_shape(arr), + vec![num_experts as i32, out, in_dim], + "stacked shape for {proj} should be [num_experts, out, in_dim]" + ); + + for expert in 0..num_experts { + let actual = mlxcel_core::slice( + arr, + &[expert as i32, 0, 0], + &[expert as i32 + 1, out, in_dim], + ); + let expected = mlxcel_core::from_slice_f32( + &vec![base + expert as f32; (out * in_dim) as usize], + &[1, out, in_dim], + ); + let close = mlxcel_core::allclose(&actual, &expected, 1e-6, 1e-6); + mlxcel_core::eval(&close); + assert!( + mlxcel_core::item_bool(&close), + "{proj} expert {expert} should preserve its source values" + ); + } + } + } + #[test] fn sanitize_weights_stacks_per_expert_gate_up_down_proj_with_scales_biases() { // Quantized per-expert gate_proj layout: weight + scales + biases per expert. diff --git a/src/models/qwen3_5_tests.rs b/src/models/qwen3_5_tests.rs index f84a176bd..e61e9af49 100644 --- a/src/models/qwen3_5_tests.rs +++ b/src/models/qwen3_5_tests.rs @@ -21,8 +21,7 @@ //! that need a real Qwen 3.5 model and are gated behind hardware availability. use super::qwen3_5::{ - Qwen35Config, rebuild_with_zero_tail, sanitize_weights, - zero_per_row_kv_tail, + Qwen35Config, rebuild_with_zero_tail, sanitize_weights, zero_per_row_kv_tail, }; use mlxcel_core::dtype; use mlxcel_core::layers::KVCache; @@ -310,34 +309,3 @@ fn sanitize_weights_drops_lm_head_when_tied_embeddings() { ); assert!(sanitized.contains_key("model.embed_tokens.weight")); } - -#[test] -#[ignore = "requires serial MLX execution"] -fn sanitize_weights_stacks_per_expert_switch_proj_names_for_loader() { - let root = "model.layers.0.mlp.experts"; - let mut weights = WeightMap::new(); - for expert in 0..2 { - for proj in ["gate_proj", "up_proj", "down_proj"] { - weights.insert( - format!("{root}.{expert}.{proj}.weight"), - mlxcel_core::from_slice_f32(&[expert as f32; 8], &[2, 4]), - ); - } - } - - let mut config = make_tiny_config(); - config.num_experts = 2; - config.num_experts_per_tok = 1; - config.decoder_sparse_step = 1; - config.moe_intermediate_size = 2; - config.shared_expert_intermediate_size = 2; - - let sanitized = sanitize_weights(weights, &config); - - assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.gate_proj.weight")); - assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.up_proj.weight")); - assert!(sanitized.contains_key("model.layers.0.mlp.switch_mlp.down_proj.weight")); - assert!(!sanitized.contains_key("model.layers.0.mlp.switch_mlp.w1.weight")); - assert!(!sanitized.contains_key("model.layers.0.mlp.switch_mlp.w2.weight")); - assert!(!sanitized.contains_key("model.layers.0.mlp.switch_mlp.w3.weight")); -}