use crate::backends::gguf::types::MetaValue;
use crate::convert::source_reader::HfTensor;
use crate::quantize::ggml_quants::SourceDtype;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MappedTensor {
Direct(String),
Drop,
}
pub fn map_tensor_name(hf_name: &str) -> Option<MappedTensor> {
if is_offpath_modality_tensor(hf_name) {
return Some(MappedTensor::Drop);
}
let working_cow = strip_language_model_prefix(hf_name);
let working: &str = working_cow.as_ref();
match working {
"model.embed_tokens.weight" => {
return Some(MappedTensor::Direct("token_embd.weight".to_string()));
}
"model.norm.weight" => {
return Some(MappedTensor::Direct("output_norm.weight".to_string()));
}
"lm_head.weight" => {
return Some(MappedTensor::Direct("output.weight".to_string()));
}
GEMMA4_ROPE_FREQS_TENSOR_NAME => {
return Some(MappedTensor::Direct(
GEMMA4_ROPE_FREQS_TENSOR_NAME.to_string(),
));
}
_ => {}
}
let stripped = working.strip_prefix("model.layers.")?;
let dot = stripped.find('.')?;
let (layer_str, rest_with_dot) = stripped.split_at(dot);
let layer: usize = layer_str.parse().ok()?;
if layer.to_string() != layer_str {
return None;
}
let rest = &rest_with_dot[1..];
let gguf_suffix: &str = match rest {
"input_layernorm.weight" => "attn_norm.weight",
"post_attention_layernorm.weight" => "post_attention_norm.weight",
"pre_feedforward_layernorm.weight" => "ffn_norm.weight",
"pre_feedforward_layernorm_2.weight" => "pre_ffw_norm_2.weight",
"post_feedforward_layernorm.weight" => "post_ffw_norm.weight",
"post_feedforward_layernorm_1.weight" => "post_ffw_norm_1.weight",
"post_feedforward_layernorm_2.weight" => "post_ffw_norm_2.weight",
"self_attn.q_proj.weight" => "attn_q.weight",
"self_attn.k_proj.weight" => "attn_k.weight",
"self_attn.v_proj.weight" => "attn_v.weight",
"self_attn.o_proj.weight" => "attn_output.weight",
"self_attn.q_norm.weight" => "attn_q_norm.weight",
"self_attn.k_norm.weight" => "attn_k_norm.weight",
"mlp.gate_proj.weight" => "ffn_gate.weight",
"mlp.up_proj.weight" => "ffn_up.weight",
"mlp.down_proj.weight" => "ffn_down.weight",
"experts.gate_up_proj" => "ffn_gate_up_exps.weight",
"experts.down_proj" => "ffn_down_exps.weight",
"router.proj.weight" => "ffn_gate_inp.weight",
"router.scale" => "ffn_gate_inp.scale",
"router.per_expert_scale" => "ffn_down_exps.scale",
"layer_scalar" => "layer_output_scale.weight",
_ => return None,
};
Some(MappedTensor::Direct(format!("blk.{layer}.{gguf_suffix}")))
}
fn strip_language_model_prefix(hf_name: &str) -> std::borrow::Cow<'_, str> {
if let Some(idx) = hf_name.find("language_model.") {
let head = &hf_name[..idx];
let tail = &hf_name[idx + "language_model.".len()..];
std::borrow::Cow::Owned(format!("{head}{tail}"))
} else {
std::borrow::Cow::Borrowed(hf_name)
}
}
fn is_offpath_modality_tensor(hf_name: &str) -> bool {
hf_name.contains("model.vision_tower.")
|| hf_name.contains("model.embed_vision.")
|| hf_name.contains("model.audio_tower.")
|| hf_name.contains("model.embed_audio.")
|| hf_name.starts_with("vision_model.")
|| hf_name.starts_with("audio_model.")
}
pub fn build_metadata(
config: &serde_json::Value,
file_type: u32,
model_card: Option<&crate::convert::model_card::ModelCard>,
sampling: Option<&crate::convert::model_card::SamplingConfig>,
model_dir_basename: Option<&str>,
) -> Vec<(String, MetaValue)> {
use crate::convert::model_card::{
emit_general_postlude, emit_general_prelude, get_model_id_components,
};
let raw_name = config
.get("_name_or_path")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.or_else(|| model_dir_basename.map(|s| s.to_string()))
.unwrap_or_else(|| "model".to_string());
let id_components = get_model_id_components(&raw_name);
let display_name = id_components
.name
.clone()
.unwrap_or_else(|| raw_name.clone());
let hidden_size = config["hidden_size"]
.as_u64()
.expect("config.json missing required key `hidden_size`") as u32;
let n_layers = config["num_hidden_layers"]
.as_u64()
.expect("config.json missing required key `num_hidden_layers`") as u32;
let ffn_len = config["intermediate_size"]
.as_u64()
.expect("config.json missing required key `intermediate_size`") as u32;
let n_head = config["num_attention_heads"]
.as_u64()
.expect("config.json missing required key `num_attention_heads`") as u32;
let ctx_len = config["max_position_embeddings"]
.as_u64()
.expect("config.json missing required key `max_position_embeddings`")
as u32;
let rms_eps = config["rms_norm_eps"]
.as_f64()
.expect("config.json missing required key `rms_norm_eps`") as f32;
let head_dim = config["head_dim"]
.as_u64()
.expect("config.json missing required key `head_dim`") as u32;
let n_head_kv = config
.get("num_key_value_heads")
.and_then(|v| v.as_u64())
.map(|x| x as u32)
.unwrap_or(n_head);
let global_head_dim = config
.get("global_head_dim")
.and_then(|v| v.as_u64())
.map(|x| x as u32)
.unwrap_or(head_dim);
let sliding_window = config
.get("sliding_window")
.and_then(|v| v.as_u64())
.map(|x| x as u32)
.unwrap_or(4096);
let rope_theta = resolve_rope_theta(config);
let num_kv_shared_layers = config["num_kv_shared_layers"]
.as_u64()
.expect("config.json missing required key `num_kv_shared_layers`")
as u32;
let layer_types_raw = config["layer_types"]
.as_array()
.expect("config.json missing required key `layer_types` (array)");
if layer_types_raw.len() as u32 != n_layers {
panic!(
"config.json `layer_types` array length {} does not match \
`num_hidden_layers` {}",
layer_types_raw.len(),
n_layers,
);
}
let sliding_window_pattern: Vec<bool> = layer_types_raw
.iter()
.map(|v| v.as_str().expect("`layer_types[i]` must be a string") == "sliding_attention")
.collect();
let hidden_size_per_layer_input = config
.get("hidden_size_per_layer_input")
.and_then(|v| v.as_u64())
.map(|x| x as u32)
.unwrap_or(0);
let partial_rotary_factor_swa = config
.get("partial_rotary_factor")
.and_then(|v| v.as_f64())
.unwrap_or(1.0);
let n_rot_full = global_head_dim;
let n_rot_swa = (head_dim as f64 * partial_rotary_factor_swa) as u32;
let mut kv: Vec<(String, MetaValue)> = emit_general_prelude(
"gemma4",
display_name,
&id_components,
None, model_card,
sampling,
);
let use_double_wide_mlp = config
.get("use_double_wide_mlp")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let feed_forward_length_kv: MetaValue = if use_double_wide_mlp {
let first_kv_shared_layer_idx = n_layers.saturating_sub(num_kv_shared_layers);
let n_ff_arr: Vec<u32> = (0..n_layers)
.map(|il| {
if il < first_kv_shared_layer_idx {
ffn_len
} else {
ffn_len * 2
}
})
.collect();
MetaValue::ArrayU32(n_ff_arr)
} else {
MetaValue::U32(ffn_len)
};
let num_global_kv = config
.get("num_global_key_value_heads")
.and_then(|v| v.as_u64())
.map(|x| x as u32);
let head_count_kv_kv: MetaValue = if let Some(num_global_kv) = num_global_kv {
if num_global_kv != n_head_kv {
let arr: Vec<i32> = sliding_window_pattern
.iter()
.map(|&is_swa| {
if is_swa {
n_head_kv as i32
} else {
num_global_kv as i32
}
})
.collect();
MetaValue::ArrayI32(arr)
} else {
MetaValue::U32(n_head_kv)
}
} else {
MetaValue::U32(n_head_kv)
};
kv.push(("gemma4.block_count".into(), MetaValue::U32(n_layers)));
kv.push(("gemma4.context_length".into(), MetaValue::U32(ctx_len)));
kv.push((
"gemma4.embedding_length".into(),
MetaValue::U32(hidden_size),
));
kv.push(("gemma4.feed_forward_length".into(), feed_forward_length_kv));
kv.push(("gemma4.attention.head_count".into(), MetaValue::U32(n_head)));
kv.push(("gemma4.attention.head_count_kv".into(), head_count_kv_kv));
kv.push(("gemma4.rope.freq_base".into(), MetaValue::F32(rope_theta)));
if let Some(swa_theta) = config
.get("rope_parameters")
.and_then(|v| v.get("sliding_attention"))
.and_then(|v| v.get("rope_theta"))
.and_then(|v| v.as_f64())
{
kv.push((
"gemma4.rope.freq_base_swa".into(),
MetaValue::F32(swa_theta as f32),
));
}
kv.push((
"gemma4.attention.layer_norm_rms_epsilon".into(),
MetaValue::F32(rms_eps),
));
if let Some(n_experts) = config.get("num_experts").and_then(|v| v.as_u64()) {
kv.push((
"gemma4.expert_count".into(),
MetaValue::U32(n_experts as u32),
));
}
if let Some(top_k) = config
.get("top_k_experts")
.or_else(|| config.get("num_experts_per_tok"))
.and_then(|v| v.as_u64())
{
kv.push((
"gemma4.expert_used_count".into(),
MetaValue::U32(top_k as u32),
));
}
kv.push((
"gemma4.attention.key_length".into(),
MetaValue::U32(global_head_dim),
));
kv.push((
"gemma4.attention.value_length".into(),
MetaValue::U32(global_head_dim),
));
if let Some(softcap) = config
.get("final_logit_softcapping")
.and_then(|v| v.as_f64())
{
kv.push((
"gemma4.final_logit_softcapping".into(),
MetaValue::F32(softcap as f32),
));
}
kv.push((
"gemma4.attention.sliding_window".into(),
MetaValue::U32(sliding_window),
));
kv.push((
"gemma4.attention.shared_kv_layers".into(),
MetaValue::U32(num_kv_shared_layers),
));
kv.push((
"gemma4.embedding_length_per_layer_input".into(),
MetaValue::U32(hidden_size_per_layer_input),
));
kv.push((
"gemma4.attention.sliding_window_pattern".into(),
MetaValue::ArrayBool(sliding_window_pattern.clone()),
));
kv.push((
"gemma4.attention.key_length_swa".into(),
MetaValue::U32(head_dim),
));
kv.push((
"gemma4.attention.value_length_swa".into(),
MetaValue::U32(head_dim),
));
if let Some(moe_ffn) = config
.get("moe_intermediate_size")
.or_else(|| config.get("expert_intermediate_size"))
.and_then(|v| v.as_u64())
{
kv.push((
"gemma4.expert_feed_forward_length".into(),
MetaValue::U32(moe_ffn as u32),
));
}
kv.push((
"gemma4.rope.dimension_count".into(),
MetaValue::U32(n_rot_full),
));
kv.push((
"gemma4.rope.dimension_count_swa".into(),
MetaValue::U32(n_rot_swa),
));
kv.extend(emit_general_postlude(file_type));
kv
}
pub const GEMMA4_ROPE_FREQS_TENSOR_NAME: &str = "rope_freqs.weight";
pub fn build_synthesized_tensors(config: &serde_json::Value) -> Vec<HfTensor> {
let head_dim = config["head_dim"]
.as_u64()
.expect("config.json missing required key `head_dim`") as u32;
let global_head_dim = config
.get("global_head_dim")
.and_then(|v| v.as_u64())
.map(|x| x as u32)
.unwrap_or(head_dim);
let rope_params_full = config
.get("rope_parameters")
.and_then(|v| v.get("full_attention"))
.expect(
"config.json missing required path \
`rope_parameters.full_attention` (required for \
ROPE_FREQS synthesis — see gemma.py:704-717)",
);
let rope_type = rope_params_full
.get("rope_type")
.and_then(|v| v.as_str())
.expect(
"config.json missing required key \
`rope_parameters.full_attention.rope_type` \
(see gemma.py:705)",
);
assert_eq!(
rope_type, "proportional",
"ROPE_FREQS synthesis only valid for rope_type=proportional, got `{rope_type}` (gemma.py:705)"
);
let partial_rotary_factor_full = rope_params_full
.get("partial_rotary_factor")
.and_then(|v| v.as_f64())
.expect(
"config.json missing required key \
`rope_parameters.full_attention.partial_rotary_factor`",
);
let table_len = (global_head_dim / 2) as usize;
let n_rot_full = (global_head_dim as f64 * partial_rotary_factor_full / 2.0) as usize;
assert!(
n_rot_full <= table_len,
"ROPE_FREQS synthesis: n_rot_full ({n_rot_full}) > table_len ({table_len}) — \
invalid combination of global_head_dim={global_head_dim} and \
partial_rotary_factor={partial_rotary_factor_full}; see gemma.py:713-715"
);
let n_unrot_full = table_len - n_rot_full;
let mut values = Vec::with_capacity(table_len);
values.extend(std::iter::repeat(1.0_f32).take(n_rot_full));
values.extend(std::iter::repeat(1.0e30_f32).take(n_unrot_full));
vec![HfTensor {
name: GEMMA4_ROPE_FREQS_TENSOR_NAME.to_string(),
shape: vec![table_len],
source_dtype: SourceDtype::F32,
data: values,
}]
}
fn resolve_rope_theta(config: &serde_json::Value) -> f32 {
if let Some(rt) = config.get("rope_theta").and_then(|v| v.as_f64()) {
return rt as f32;
}
if let Some(rp) = config.get("rope_parameters") {
if let Some(rt) = rp
.get("full_attention")
.and_then(|v| v.get("rope_theta"))
.and_then(|v| v.as_f64())
{
return rt as f32;
}
if let Some(rt) = rp.get("rope_theta").and_then(|v| v.as_f64()) {
return rt as f32;
}
}
10000.0
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn gemma4_strips_language_model_prefix() {
for hf in [
"model.embed_tokens.weight",
"model.language_model.embed_tokens.weight",
] {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct("token_embd.weight".into())),
"{hf:?} must map to token_embd.weight"
);
}
for hf in [
"model.layers.7.self_attn.q_proj.weight",
"model.language_model.layers.7.self_attn.q_proj.weight",
] {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct("blk.7.attn_q.weight".into())),
"{hf:?} must map to blk.7.attn_q.weight"
);
}
for hf in ["model.norm.weight", "model.language_model.norm.weight"] {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct("output_norm.weight".into())),
);
}
}
#[test]
fn gemma4_experts_gate_up_proj_returns_fused_direct() {
let cases = [
(
"model.language_model.layers.0.experts.gate_up_proj",
"blk.0.ffn_gate_up_exps.weight",
),
(
"model.layers.0.experts.gate_up_proj",
"blk.0.ffn_gate_up_exps.weight",
),
(
"model.language_model.layers.15.experts.gate_up_proj",
"blk.15.ffn_gate_up_exps.weight",
),
(
"model.language_model.layers.29.experts.gate_up_proj",
"blk.29.ffn_gate_up_exps.weight",
),
];
for (hf, gguf) in cases {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct(gguf.into())),
"{hf:?} → expected {gguf:?}"
);
}
}
#[test]
fn gemma4_experts_down_proj_maps_to_ffn_down_exps() {
for hf in [
"model.language_model.layers.0.experts.down_proj",
"model.layers.0.experts.down_proj",
] {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct("blk.0.ffn_down_exps.weight".into())),
"{hf:?} → expected blk.0.ffn_down_exps.weight"
);
}
}
#[test]
fn gemma4_dual_norms_disambiguated() {
let pairs: &[(&str, &str)] = &[
(
"model.language_model.layers.0.pre_feedforward_layernorm.weight",
"blk.0.ffn_norm.weight",
),
(
"model.language_model.layers.0.pre_feedforward_layernorm_2.weight",
"blk.0.pre_ffw_norm_2.weight",
),
(
"model.language_model.layers.0.post_feedforward_layernorm.weight",
"blk.0.post_ffw_norm.weight",
),
(
"model.language_model.layers.0.post_feedforward_layernorm_1.weight",
"blk.0.post_ffw_norm_1.weight",
),
(
"model.language_model.layers.0.post_feedforward_layernorm_2.weight",
"blk.0.post_ffw_norm_2.weight",
),
];
let mut seen_gguf: std::collections::HashSet<String> = std::collections::HashSet::new();
for (hf, gguf) in pairs {
let got = map_tensor_name(hf);
assert_eq!(
got,
Some(MappedTensor::Direct((*gguf).into())),
"{hf:?} → expected {gguf:?}, got {got:?}"
);
assert!(
seen_gguf.insert((*gguf).into()),
"norm GGUF target {gguf:?} appears twice (load-bearing for correctness)"
);
}
}
#[test]
fn gemma4_router_proj() {
for (hf, gguf) in [
(
"model.language_model.layers.0.router.proj.weight",
"blk.0.ffn_gate_inp.weight",
),
(
"model.layers.13.router.proj.weight",
"blk.13.ffn_gate_inp.weight",
),
] {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct(gguf.into())),
"{hf:?} → expected {gguf:?}"
);
}
}
#[test]
fn gemma4_router_scale_subnames() {
assert_eq!(
map_tensor_name("model.language_model.layers.0.router.scale"),
Some(MappedTensor::Direct("blk.0.ffn_gate_inp.scale".into()))
);
assert_eq!(
map_tensor_name("model.language_model.layers.7.router.per_expert_scale"),
Some(MappedTensor::Direct("blk.7.ffn_down_exps.scale".into()))
);
}
#[test]
fn gemma4_layer_scalar() {
for layer in [0usize, 15, 29] {
let hf = format!("model.language_model.layers.{layer}.layer_scalar");
let want = format!("blk.{layer}.layer_output_scale.weight");
assert_eq!(
map_tensor_name(&hf),
Some(MappedTensor::Direct(want.clone())),
"{hf:?} → expected {want:?}"
);
}
}
#[test]
fn gemma4_parallel_dense_ffn() {
let cases = [
(
"model.language_model.layers.3.mlp.gate_proj.weight",
"blk.3.ffn_gate.weight",
),
(
"model.language_model.layers.3.mlp.up_proj.weight",
"blk.3.ffn_up.weight",
),
(
"model.language_model.layers.3.mlp.down_proj.weight",
"blk.3.ffn_down.weight",
),
];
for (hf, gguf) in cases {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct(gguf.into())),
"{hf:?} → expected {gguf:?}"
);
}
}
#[test]
fn gemma4_attention_block() {
let cases = [
(
"model.language_model.layers.5.self_attn.q_proj.weight",
"blk.5.attn_q.weight",
),
(
"model.language_model.layers.5.self_attn.k_proj.weight",
"blk.5.attn_k.weight",
),
(
"model.language_model.layers.5.self_attn.v_proj.weight",
"blk.5.attn_v.weight",
),
(
"model.language_model.layers.5.self_attn.o_proj.weight",
"blk.5.attn_output.weight",
),
(
"model.language_model.layers.5.self_attn.q_norm.weight",
"blk.5.attn_q_norm.weight",
),
(
"model.language_model.layers.5.self_attn.k_norm.weight",
"blk.5.attn_k_norm.weight",
),
];
for (hf, gguf) in cases {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Direct(gguf.into())),
"{hf:?} → expected {gguf:?}"
);
}
}
#[test]
fn gemma4_vision_tower_returns_drop() {
let cases = [
"model.vision_tower.encoder.layers.0.self_attn.q_proj.linear.weight",
"model.vision_tower.patch_embedder.input_proj.weight",
"model.vision_tower.patch_embedder.position_embedding_table",
"model.vision_tower.std_bias",
"model.vision_tower.std_scale",
"model.embed_vision.embedding_projection.weight",
"model.audio_tower.encoder.layers.0.self_attn.q_proj.weight",
"model.embed_audio.linear.weight",
"vision_model.encoder.layers.0.attn.q_proj.weight",
"audio_model.encoder.layers.0.attn.q_proj.weight",
];
for hf in cases {
assert_eq!(
map_tensor_name(hf),
Some(MappedTensor::Drop),
"{hf:?} → expected Drop (off-path modality)"
);
}
}
#[test]
fn gemma4_tensor_name_rejects_unknown_kinds() {
let cases = [
"model.unknown.weight",
"transformer.layers.0.attn.weight",
"model.language_model.layers.0.self_attn.q_proj.bias",
"model.language_model.layers.01.self_attn.q_proj.weight",
"model.language_model.layers..self_attn.q_proj.weight",
"model.language_model.layers.self_attn.q_proj.weight",
"model.language_model.layers.0.unknown.weight",
"model.language_model.layers.0.mlp.shared_expert.gate_proj.weight",
"model.language_model.layers.0.mlp.experts.0.gate_proj.weight",
];
for hf in cases {
assert_eq!(
map_tensor_name(hf),
None,
"{hf:?} → expected None (unmapped)"
);
}
}
#[test]
fn gemma4_metadata_from_real_config() {
let layer_types = (0..30u32)
.map(|i| {
if (i + 1) % 6 == 0 {
"full_attention"
} else {
"sliding_attention"
}
})
.collect::<Vec<_>>();
let cfg = json!({
"_name_or_path": "google/gemma-4-26b-a4b-it",
"model_type": "gemma4_text",
"hidden_size": 2816,
"num_hidden_layers": 30,
"intermediate_size": 2112,
"moe_intermediate_size": 704,
"num_attention_heads": 16,
"num_key_value_heads": 8,
"num_global_key_value_heads": 2,
"num_kv_shared_layers": 0,
"hidden_size_per_layer_input": 0,
"layer_types": layer_types,
"use_double_wide_mlp": false,
"head_dim": 256,
"global_head_dim": 512,
"max_position_embeddings": 262144,
"rms_norm_eps": 1.0e-6,
"sliding_window": 1024,
"num_experts": 128,
"top_k_experts": 4,
"rope_parameters": {
"full_attention": {
"rope_theta": 1_000_000.0,
"rope_type": "proportional",
"partial_rotary_factor": 0.25,
}
},
});
let kv = build_metadata(&cfg, 17 , None, None, None);
let by_key: std::collections::HashMap<_, _> =
kv.iter().map(|(k, v)| (k.as_str(), v.clone())).collect();
assert_eq!(
by_key["general.architecture"],
MetaValue::String("gemma4".into()),
"Gemma 4 emits general.architecture=gemma4 (LLM_ARCH_GEMMA4)"
);
assert_eq!(
by_key["general.name"],
MetaValue::String("Gemma 4 26b A4B It".into())
);
assert_eq!(by_key["gemma4.context_length"], MetaValue::U32(262144));
assert_eq!(by_key["gemma4.embedding_length"], MetaValue::U32(2816));
assert_eq!(by_key["gemma4.block_count"], MetaValue::U32(30));
assert_eq!(by_key["gemma4.attention.head_count"], MetaValue::U32(16));
assert_eq!(by_key["gemma4.attention.key_length"], MetaValue::U32(512));
assert_eq!(by_key["gemma4.attention.value_length"], MetaValue::U32(512));
assert_eq!(
by_key["gemma4.attention.key_length_swa"],
MetaValue::U32(256)
);
assert_eq!(
by_key["gemma4.attention.value_length_swa"],
MetaValue::U32(256)
);
assert_eq!(
by_key["gemma4.attention.layer_norm_rms_epsilon"],
MetaValue::F32(1.0e-6)
);
assert_eq!(
by_key["gemma4.rope.freq_base"],
MetaValue::F32(1_000_000.0),
"rope_theta must come from rope_parameters.full_attention.rope_theta"
);
assert_eq!(
by_key["gemma4.attention.sliding_window"],
MetaValue::U32(1024)
);
assert_eq!(by_key["gemma4.expert_count"], MetaValue::U32(128));
assert_eq!(by_key["gemma4.expert_used_count"], MetaValue::U32(4));
assert_eq!(
by_key["gemma4.expert_feed_forward_length"],
MetaValue::U32(704)
);
assert_eq!(by_key["general.file_type"], MetaValue::U32(17));
assert_eq!(
by_key["gemma4.attention.shared_kv_layers"],
MetaValue::U32(0),
"shared_kv_layers from num_kv_shared_layers"
);
assert_eq!(
by_key["gemma4.embedding_length_per_layer_input"],
MetaValue::U32(0),
"embedding_length_per_layer_input from hidden_size_per_layer_input"
);
let swa_pattern = match &by_key["gemma4.attention.sliding_window_pattern"] {
MetaValue::ArrayBool(v) => v.clone(),
other => panic!("expected ArrayBool, got {other:?}"),
};
assert_eq!(
swa_pattern.len(),
30,
"swa pattern must have one entry per layer"
);
for (i, &is_swa) in swa_pattern.iter().enumerate() {
let expected_swa = !matches!(i, 5 | 11 | 17 | 23 | 29);
assert_eq!(is_swa, expected_swa, "layer {i} swa-flag mismatch");
}
assert_eq!(
by_key["gemma4.rope.dimension_count"],
MetaValue::U32(512),
"rope.dimension_count = int(global_head_dim)"
);
assert_eq!(
by_key["gemma4.rope.dimension_count_swa"],
MetaValue::U32(256),
"rope.dimension_count_swa = int(head_dim * partial_rotary_factor=1.0)"
);
let hck = match &by_key["gemma4.attention.head_count_kv"] {
MetaValue::ArrayI32(v) => v.clone(),
other => {
panic!("expected ArrayI32 for head_count_kv (global=2 ≠ swa=8), got {other:?}")
}
};
assert_eq!(hck.len(), 30, "head_count_kv array length = block_count");
for (i, &hck_i) in hck.iter().enumerate() {
let expected = if matches!(i, 5 | 11 | 17 | 23 | 29) {
2
} else {
8
};
assert_eq!(hck_i, expected, "layer {i} head_count_kv mismatch");
}
assert_eq!(
by_key["gemma4.feed_forward_length"],
MetaValue::U32(2112),
"feed_forward_length scalar when use_double_wide_mlp=false"
);
}
#[test]
fn gemma4_metadata_optional_key_defaults() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 1,
"intermediate_size": 64,
"num_attention_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 0,
"layer_types": ["sliding_attention"],
});
let kv = build_metadata(&cfg, 0, None, None, None);
let by_key: std::collections::HashMap<_, _> =
kv.iter().map(|(k, v)| (k.as_str(), v.clone())).collect();
assert_eq!(by_key["general.name"], MetaValue::String("Model".into()),);
assert_eq!(
by_key["gemma4.attention.head_count_kv"],
MetaValue::U32(4),
"num_key_value_heads defaults to num_attention_heads"
);
assert_eq!(
by_key["gemma4.attention.key_length"],
MetaValue::U32(16),
"global_head_dim defaults to head_dim"
);
assert_eq!(
by_key["gemma4.rope.freq_base"],
MetaValue::F32(10000.0),
"rope_theta defaults to 10000.0"
);
assert_eq!(
by_key["gemma4.attention.sliding_window"],
MetaValue::U32(4096),
"sliding_window defaults to 4096"
);
assert_eq!(
by_key["general.architecture"],
MetaValue::String("gemma4".into()),
);
assert!(
!by_key.contains_key("gemma4.expert_count"),
"expert_count must be absent when num_experts is absent"
);
assert!(
!by_key.contains_key("gemma4.expert_used_count"),
"expert_used_count must be absent when top_k_experts is absent"
);
assert!(
!by_key.contains_key("gemma4.expert_feed_forward_length"),
"expert_feed_forward_length must be absent when moe_intermediate_size is absent"
);
assert_eq!(by_key["general.file_type"], MetaValue::U32(0));
assert_eq!(
by_key["gemma4.attention.shared_kv_layers"],
MetaValue::U32(0),
);
assert_eq!(
by_key["gemma4.embedding_length_per_layer_input"],
MetaValue::U32(0),
);
assert_eq!(
by_key["gemma4.attention.sliding_window_pattern"],
MetaValue::ArrayBool(vec![true]),
);
assert_eq!(by_key["gemma4.rope.dimension_count"], MetaValue::U32(16));
assert_eq!(
by_key["gemma4.rope.dimension_count_swa"],
MetaValue::U32(16)
);
assert_eq!(
by_key["gemma4.attention.head_count_kv"],
MetaValue::U32(4),
"head_count_kv stays scalar when num_global_key_value_heads is absent"
);
assert_eq!(
by_key["gemma4.feed_forward_length"],
MetaValue::U32(64),
"feed_forward_length stays scalar when use_double_wide_mlp is false"
);
}
#[test]
fn gemma4_resolve_rope_theta_three_paths() {
let cfg = json!({ "rope_theta": 500_000.0 });
assert_eq!(resolve_rope_theta(&cfg), 500_000.0);
let cfg = json!({
"rope_parameters": {
"full_attention": {
"rope_theta": 1_000_000.0,
}
}
});
assert_eq!(resolve_rope_theta(&cfg), 1_000_000.0);
let cfg = json!({
"rope_parameters": { "rope_theta": 250_000.0 }
});
assert_eq!(resolve_rope_theta(&cfg), 250_000.0);
let cfg = json!({});
assert_eq!(resolve_rope_theta(&cfg), 10000.0);
}
#[test]
fn gemma4_metadata_ftype_round_trips() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 1,
"intermediate_size": 64,
"num_attention_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 0,
"layer_types": ["sliding_attention"],
});
for &ftype in &[0u32, 1, 7, 15, 17, 23] {
let kv = build_metadata(&cfg, ftype, None, None, None);
let by_key: std::collections::HashMap<_, _> =
kv.iter().map(|(k, v)| (k.as_str(), v.clone())).collect();
assert_eq!(
by_key["general.file_type"],
MetaValue::U32(ftype),
"file_type {ftype} must round-trip as MetaValue::U32"
);
}
}
#[test]
fn gemma4_strip_language_model_helper() {
assert_eq!(
strip_language_model_prefix("model.language_model.embed_tokens.weight").as_ref(),
"model.embed_tokens.weight"
);
assert_eq!(
strip_language_model_prefix("model.embed_tokens.weight").as_ref(),
"model.embed_tokens.weight"
);
assert_eq!(
strip_language_model_prefix("outer.language_model.inner").as_ref(),
"outer.inner"
);
}
#[test]
#[should_panic(expected = "num_kv_shared_layers")]
fn gemma4_metadata_panics_on_missing_num_kv_shared_layers() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 1,
"intermediate_size": 64,
"num_attention_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"layer_types": ["sliding_attention"],
});
let _ = build_metadata(&cfg, 0, None, None, None);
}
#[test]
#[should_panic(expected = "layer_types")]
fn gemma4_metadata_panics_on_missing_layer_types() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 1,
"intermediate_size": 64,
"num_attention_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 0,
});
let _ = build_metadata(&cfg, 0, None, None, None);
}
#[test]
#[should_panic(expected = "layer_types")]
fn gemma4_metadata_panics_on_layer_types_length_mismatch() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 3,
"intermediate_size": 64,
"num_attention_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 0,
"layer_types": ["sliding_attention", "full_attention"], });
let _ = build_metadata(&cfg, 0, None, None, None);
}
#[test]
fn gemma4_metadata_feed_forward_length_array_on_dwm() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 6,
"intermediate_size": 64,
"num_attention_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 2,
"layer_types": [
"sliding_attention",
"sliding_attention",
"sliding_attention",
"sliding_attention",
"sliding_attention",
"sliding_attention",
],
"use_double_wide_mlp": true,
});
let kv = build_metadata(&cfg, 0, None, None, None);
let by_key: std::collections::HashMap<_, _> =
kv.iter().map(|(k, v)| (k.as_str(), v.clone())).collect();
assert_eq!(
by_key["gemma4.feed_forward_length"],
MetaValue::ArrayU32(vec![64, 64, 64, 64, 128, 128]),
"dwm array: last num_kv_shared_layers entries are 2*n_ff"
);
}
#[test]
fn gemma4_metadata_head_count_kv_array_path() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 4,
"intermediate_size": 64,
"num_attention_heads": 8,
"num_key_value_heads": 4,
"num_global_key_value_heads": 2, "head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 0,
"layer_types": [
"sliding_attention",
"sliding_attention",
"full_attention",
"full_attention",
],
});
let kv = build_metadata(&cfg, 0, None, None, None);
let by_key: std::collections::HashMap<_, _> =
kv.iter().map(|(k, v)| (k.as_str(), v.clone())).collect();
assert_eq!(
by_key["gemma4.attention.head_count_kv"],
MetaValue::ArrayI32(vec![4, 4, 2, 2]),
);
}
#[test]
fn gemma4_metadata_head_count_kv_scalar_when_equal() {
let cfg = json!({
"hidden_size": 32,
"num_hidden_layers": 2,
"intermediate_size": 64,
"num_attention_heads": 4,
"num_key_value_heads": 4,
"num_global_key_value_heads": 4,
"head_dim": 16,
"max_position_embeddings": 2048,
"rms_norm_eps": 1.0e-6,
"num_kv_shared_layers": 0,
"layer_types": ["sliding_attention", "full_attention"],
});
let kv = build_metadata(&cfg, 0, None, None, None);
let by_key: std::collections::HashMap<_, _> =
kv.iter().map(|(k, v)| (k.as_str(), v.clone())).collect();
assert_eq!(
by_key["gemma4.attention.head_count_kv"],
MetaValue::U32(4),
"equal global/swa kv heads → scalar form"
);
}
#[test]
fn gemma4_synthesized_rope_freqs_real_config() {
let cfg = json!({
"head_dim": 256,
"global_head_dim": 512,
"rope_parameters": {
"full_attention": {
"partial_rotary_factor": 0.25,
"rope_type": "proportional",
}
},
});
let tensors = build_synthesized_tensors(&cfg);
assert_eq!(
tensors.len(),
1,
"exactly one synthesized tensor for Gemma 4"
);
let t = &tensors[0];
assert_eq!(t.name, "rope_freqs.weight");
assert_eq!(t.shape, vec![256]);
assert_eq!(t.data.len(), 256);
for (i, &v) in t.data.iter().enumerate() {
if i < 64 {
assert_eq!(v, 1.0, "entry {i} (rotated dim) must be 1.0");
} else {
assert_eq!(v, 1.0e30, "entry {i} (unrotated dim) must be 1e30");
}
}
}
#[test]
#[should_panic(expected = "rope_parameters.full_attention")]
fn gemma4_synthesized_rope_freqs_panics_on_missing_rope_params() {
let cfg = json!({
"head_dim": 256,
"global_head_dim": 512,
});
let _ = build_synthesized_tensors(&cfg);
}
#[test]
#[should_panic(expected = "partial_rotary_factor")]
fn gemma4_synthesized_rope_freqs_panics_on_missing_partial_rotary_factor() {
let cfg = json!({
"head_dim": 256,
"global_head_dim": 512,
"rope_parameters": {
"full_attention": {
"rope_type": "proportional",
}
},
});
let _ = build_synthesized_tensors(&cfg);
}
#[test]
#[should_panic(expected = "rope_parameters.full_attention.rope_type")]
fn gemma4_synthesized_rope_freqs_panics_on_missing_rope_type() {
let cfg = json!({
"head_dim": 256,
"global_head_dim": 512,
"rope_parameters": {
"full_attention": {
"partial_rotary_factor": 0.25,
}
},
});
let _ = build_synthesized_tensors(&cfg);
}
#[test]
#[should_panic(expected = "rope_type=proportional")]
fn gemma4_synthesized_rope_freqs_panics_on_wrong_rope_type() {
let cfg = json!({
"head_dim": 256,
"global_head_dim": 512,
"rope_parameters": {
"full_attention": {
"rope_type": "linear",
"partial_rotary_factor": 0.25,
}
},
});
let _ = build_synthesized_tensors(&cfg);
}
#[test]
#[should_panic(expected = "n_rot_full")]
fn gemma4_synthesized_rope_freqs_panics_on_invalid_partial_factor() {
let cfg = json!({
"head_dim": 256,
"global_head_dim": 512,
"rope_parameters": {
"full_attention": {
"rope_type": "proportional",
"partial_rotary_factor": 1.5, }
},
});
let _ = build_synthesized_tensors(&cfg);
}
#[test]
fn gemma4_maps_synthesized_rope_freqs() {
assert_eq!(
map_tensor_name("rope_freqs.weight"),
Some(MappedTensor::Direct("rope_freqs.weight".into()))
);
}
}