use super::super::tensor_ref::ArchName;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum MoeTensorRole {
RoutedExpert,
SharedExpert,
Attention,
TokenEmbd,
Output,
RouterGate,
Norm,
Ssm,
Other,
}
pub const SUPPORTED_APEX_ARCHES: &[&str] = &["qwen3moe", "qwen35moe", "gemma4", "minimax-m2"];
pub fn classify_moe_tensor(arch: ArchName, name: &str) -> MoeTensorRole {
let _ = arch;
if name == "token_embd.weight" || name == "per_layer_token_embd.weight" {
return MoeTensorRole::TokenEmbd;
}
if name == "output.weight" || name == "lm_head.weight" {
return MoeTensorRole::Output;
}
if name == "output_norm.weight" {
return MoeTensorRole::Norm;
}
if name.contains("_norm.weight") || name.ends_with(".norm.weight") {
return MoeTensorRole::Norm;
}
if name.contains("ffn_gate_inp") {
return MoeTensorRole::RouterGate;
}
if name.contains("_shexp.weight") || name.contains("_shexp_") {
return MoeTensorRole::SharedExpert;
}
if name.contains("_exps.weight") || name.contains("_exps_") {
return MoeTensorRole::RoutedExpert;
}
if name.contains("attn_qkv")
|| name.contains("attn_kv_b")
|| name.contains("attn_q.weight")
|| name.contains("attn_k.weight")
|| name.contains("attn_v.weight")
|| name.contains("attn_output")
|| name.contains("attn_gate.weight")
|| name.contains("attn_q_norm") || name.contains("attn_k_norm")
{
if name.contains("_norm.weight") {
return MoeTensorRole::Norm;
}
return MoeTensorRole::Attention;
}
if name.contains("ssm_alpha") || name.contains("ssm_beta") || name.contains("ssm_out") {
return MoeTensorRole::Ssm;
}
MoeTensorRole::Other
}
pub const fn is_apex_supported_arch(arch: ArchName) -> bool {
matches!(
arch,
ArchName::Qwen35Moe | ArchName::Qwen35MoeFull | ArchName::Gemma4 | ArchName::MiniMaxM2
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classify_routed_expert() {
for suffix in [
"ffn_gate_exps.weight",
"ffn_up_exps.weight",
"ffn_down_exps.weight",
] {
let name = format!("blk.5.{suffix}");
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, &name),
MoeTensorRole::RoutedExpert,
"{name}"
);
}
}
#[test]
fn classify_shared_expert() {
for suffix in [
"ffn_gate_shexp.weight",
"ffn_up_shexp.weight",
"ffn_down_shexp.weight",
] {
let name = format!("blk.5.{suffix}");
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, &name),
MoeTensorRole::SharedExpert,
"{name}"
);
}
}
#[test]
fn classify_ffn_gate_inp_shexp_as_router_codex_867dba20() {
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, "blk.3.ffn_gate_inp_shexp.weight"),
MoeTensorRole::RouterGate,
"ffn_gate_inp_shexp.weight is a router gate (QWEN3_NEXT), not a shared expert"
);
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, "blk.3.ffn_gate_inp.weight"),
MoeTensorRole::RouterGate,
);
}
#[test]
fn classify_attention() {
for suffix in [
"attn_q.weight",
"attn_k.weight",
"attn_v.weight",
"attn_output.weight",
"attn_qkv.weight",
] {
let name = format!("blk.0.{suffix}");
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, &name),
MoeTensorRole::Attention,
"{name}"
);
}
}
#[test]
fn classify_token_embd_and_output() {
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, "token_embd.weight"),
MoeTensorRole::TokenEmbd
);
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, "output.weight"),
MoeTensorRole::Output
);
}
#[test]
fn classify_norms_as_norm() {
for name in [
"blk.0.attn_norm.weight",
"blk.0.ffn_norm.weight",
"blk.0.attn_q_norm.weight",
"output_norm.weight",
] {
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, name),
MoeTensorRole::Norm,
"{name}"
);
}
}
#[test]
fn classify_router_gate() {
assert_eq!(
classify_moe_tensor(ArchName::Qwen35Moe, "blk.0.ffn_gate_inp.weight"),
MoeTensorRole::RouterGate
);
}
#[test]
fn supported_arches_set() {
assert!(is_apex_supported_arch(ArchName::Qwen35Moe));
assert!(is_apex_supported_arch(ArchName::Gemma4));
assert!(is_apex_supported_arch(ArchName::MiniMaxM2));
assert!(!is_apex_supported_arch(ArchName::Llama3));
assert!(!is_apex_supported_arch(ArchName::Bert));
assert!(!is_apex_supported_arch(ArchName::NomicBert));
assert!(!is_apex_supported_arch(ArchName::Qwen3VlText));
assert!(!is_apex_supported_arch(ArchName::Falcon));
}
}