pub const GLOBAL_TENSORS: &[&str] = &["token_embd.weight", "output.weight", "output_norm.weight"];
pub const LINEAR_LAYER_TENSOR_SUFFIXES: &[&str] = &[
"attn_norm.weight",
"attn_qkv.weight",
"attn_gate.weight",
"ssm_conv1d.weight",
"ssm_dt.bias",
"ssm_a", "ssm_alpha.weight",
"ssm_beta.weight",
"ssm_norm.weight",
"ssm_out.weight",
"post_attention_norm.weight",
"ffn_gate_inp.weight",
"ffn_gate_exps.weight",
"ffn_up_exps.weight",
"ffn_down_exps.weight",
"ffn_gate_inp_shexp.weight",
"ffn_gate_shexp.weight",
"ffn_up_shexp.weight",
"ffn_down_shexp.weight",
];
pub const FULL_LAYER_TENSOR_SUFFIXES: &[&str] = &[
"attn_norm.weight",
"attn_q.weight",
"attn_k.weight",
"attn_v.weight",
"attn_q_norm.weight",
"attn_k_norm.weight",
"attn_gate.weight",
"attn_output.weight",
"post_attention_norm.weight",
"ffn_gate_inp.weight",
"ffn_gate_exps.weight",
"ffn_up_exps.weight",
"ffn_down_exps.weight",
"ffn_gate_inp_shexp.weight",
"ffn_gate_shexp.weight",
"ffn_up_shexp.weight",
"ffn_down_shexp.weight",
];
pub fn tensor_names_for_layer(layer_idx: u32, kind: super::Qwen35LayerKind) -> Vec<String> {
let suffixes = match kind {
super::Qwen35LayerKind::LinearAttention => LINEAR_LAYER_TENSOR_SUFFIXES,
super::Qwen35LayerKind::FullAttention => FULL_LAYER_TENSOR_SUFFIXES,
};
suffixes
.iter()
.map(|s| format!("blk.{layer_idx}.{s}"))
.collect()
}
pub fn expected_tensor_count(cfg: &super::Qwen35Config) -> usize {
let mut n = GLOBAL_TENSORS.len();
for kind in &cfg.layer_types {
n += match kind {
super::Qwen35LayerKind::LinearAttention => LINEAR_LAYER_TENSOR_SUFFIXES.len(),
super::Qwen35LayerKind::FullAttention => FULL_LAYER_TENSOR_SUFFIXES.len(),
};
}
n
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::models::qwen35::Qwen35LayerKind;
#[test]
fn linear_layer_names_start_with_blk_prefix() {
let names = tensor_names_for_layer(3, Qwen35LayerKind::LinearAttention);
assert!(names.iter().all(|n| n.starts_with("blk.3.")));
assert!(names.iter().any(|n| n == "blk.3.attn_qkv.weight"));
assert!(names.iter().any(|n| n == "blk.3.ssm_a"));
assert!(names.iter().any(|n| n == "blk.3.ssm_out.weight"));
}
#[test]
fn full_layer_names_include_split_qkv() {
let names = tensor_names_for_layer(11, Qwen35LayerKind::FullAttention);
assert!(names.iter().any(|n| n == "blk.11.attn_q.weight"));
assert!(names.iter().any(|n| n == "blk.11.attn_k.weight"));
assert!(names.iter().any(|n| n == "blk.11.attn_v.weight"));
assert!(names.iter().any(|n| n == "blk.11.attn_q_norm.weight"));
assert!(names.iter().any(|n| n == "blk.11.attn_k_norm.weight"));
assert!(names.iter().any(|n| n == "blk.11.attn_gate.weight"));
}
#[test]
fn full_layer_has_no_ssm_tensors() {
let names = tensor_names_for_layer(7, Qwen35LayerKind::FullAttention);
assert!(!names.iter().any(|n| n.contains(".ssm_")));
assert!(!names.iter().any(|n| n.ends_with(".attn_qkv.weight")));
}
#[test]
fn linear_layer_has_no_split_qkv() {
let names = tensor_names_for_layer(0, Qwen35LayerKind::LinearAttention);
assert!(!names.iter().any(|n| n.ends_with(".attn_q.weight")));
assert!(!names.iter().any(|n| n.ends_with(".attn_k.weight")));
assert!(!names.iter().any(|n| n.ends_with(".attn_v.weight")));
assert!(names.iter().any(|n| n.ends_with(".attn_qkv.weight")));
}
#[test]
fn global_tensors_have_three() {
assert_eq!(GLOBAL_TENSORS.len(), 3);
assert!(GLOBAL_TENSORS.contains(&"token_embd.weight"));
assert!(GLOBAL_TENSORS.contains(&"output.weight"));
assert!(GLOBAL_TENSORS.contains(&"output_norm.weight"));
}
}