use super::moe::{FULL_LAYER_TENSOR_SUFFIXES, LINEAR_LAYER_TENSOR_SUFFIXES};
pub const DENSE_FFN_TENSOR_SUFFIXES: &[&str] = &[
"ffn_gate.weight", "ffn_up.weight", "ffn_down.weight", ];
pub fn dense_layer_tensor_suffixes(kind: super::Qwen35LayerKind) -> Vec<String> {
let base = match kind {
super::Qwen35LayerKind::LinearAttention => LINEAR_LAYER_TENSOR_SUFFIXES,
super::Qwen35LayerKind::FullAttention => FULL_LAYER_TENSOR_SUFFIXES,
};
let non_moe: Vec<String> = base
.iter()
.filter(|s| {
!(s.contains("_exps.weight")
|| s.contains("_shexp.weight")
|| s.ends_with("ffn_gate_inp.weight"))
})
.map(|s| s.to_string())
.collect();
let mut out = non_moe;
for ffn in DENSE_FFN_TENSOR_SUFFIXES {
out.push(ffn.to_string());
}
out
}
pub fn tensor_names_for_layer(layer_idx: u32, kind: super::Qwen35LayerKind) -> Vec<String> {
dense_layer_tensor_suffixes(kind)
.into_iter()
.map(|s| format!("blk.{layer_idx}.{s}"))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::models::qwen35::Qwen35LayerKind;
#[test]
fn dense_ffn_has_swiglu_not_moe() {
assert_eq!(DENSE_FFN_TENSOR_SUFFIXES.len(), 3);
assert!(DENSE_FFN_TENSOR_SUFFIXES.contains(&"ffn_gate.weight"));
assert!(DENSE_FFN_TENSOR_SUFFIXES.contains(&"ffn_up.weight"));
assert!(DENSE_FFN_TENSOR_SUFFIXES.contains(&"ffn_down.weight"));
}
#[test]
fn dense_full_layer_excludes_moe_ffn_tensors() {
let names = tensor_names_for_layer(3, Qwen35LayerKind::FullAttention);
assert!(!names.iter().any(|n| n.contains("_exps.weight")));
assert!(!names.iter().any(|n| n.contains("_shexp.weight")));
assert!(!names.iter().any(|n| n.ends_with(".ffn_gate_inp.weight")));
assert!(names.iter().any(|n| n == "blk.3.ffn_gate.weight"));
assert!(names.iter().any(|n| n == "blk.3.ffn_up.weight"));
assert!(names.iter().any(|n| n == "blk.3.ffn_down.weight"));
}
#[test]
fn dense_linear_layer_keeps_ssm_tensors() {
let names = tensor_names_for_layer(0, Qwen35LayerKind::LinearAttention);
assert!(names.iter().any(|n| n == "blk.0.ssm_a"));
assert!(names.iter().any(|n| n == "blk.0.ssm_out.weight"));
assert!(names.iter().any(|n| n == "blk.0.attn_qkv.weight"));
assert!(!names.iter().any(|n| n.contains("_exps.weight")));
}
}