use crate::backends::gguf::types::MetaValue;
fn strip_bert_prefix(name: &str) -> &str {
name.strip_prefix("bert.").unwrap_or(name)
}
pub fn map_tensor_name(hf_name: &str) -> Option<String> {
let name = strip_bert_prefix(hf_name);
match name {
"embeddings.word_embeddings.weight" => {
return Some("token_embd.weight".to_string());
}
"embeddings.position_embeddings.weight" => {
return Some("position_embd.weight".to_string());
}
"embeddings.token_type_embeddings.weight" => {
return Some("token_types.weight".to_string());
}
"embeddings.LayerNorm.weight" => {
return Some("token_embd_norm.weight".to_string());
}
"embeddings.LayerNorm.bias" => {
return Some("token_embd_norm.bias".to_string());
}
"pooler.dense.weight" => return Some("cls.weight".to_string()),
"pooler.dense.bias" => return Some("cls.bias".to_string()),
_ => {}
}
let stripped = name.strip_prefix("encoder.layer.")?;
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 (head, suffix) = if let Some(stem) = rest.strip_suffix(".weight") {
(stem, ".weight")
} else if let Some(stem) = rest.strip_suffix(".bias") {
(stem, ".bias")
} else {
return None;
};
let local = match head {
"attention.self.query" => "attn_q",
"attention.self.key" => "attn_k",
"attention.self.value" => "attn_v",
"attention.output.dense" => "attn_output",
"attention.output.LayerNorm" => "attn_output_norm",
"intermediate.dense" => "ffn_up",
"output.dense" => "ffn_down",
"output.LayerNorm" => "layer_output_norm",
_ => return None,
};
Some(format!("blk.{layer}.{local}{suffix}"))
}
fn pooling_type_u32(mode: Option<&str>) -> Option<u32> {
match mode {
None => Some(1), Some("mean") | Some("MEAN") => Some(1),
Some("cls") | Some("CLS") => Some(2),
Some("last") | Some("lasttoken") | Some("LAST") => Some(3),
Some("none") | Some("NONE") => Some(0),
Some("rank") | Some("RANK") => Some(4),
_ => None,
}
}
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>,
pooling_override: Option<u32>,
) -> Vec<(String, MetaValue)> {
use crate::convert::model_card::{
emit_general_postlude, emit_general_prelude, get_model_id_components,
};
let raw_name = model_dir_basename
.map(|s| s.to_string())
.or_else(|| {
config
.get("_name_or_path")
.and_then(|v| v.as_str())
.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 ln_eps = config["layer_norm_eps"]
.as_f64()
.expect("config.json missing required key `layer_norm_eps`") as f32;
let pooling_u32 = pooling_override.unwrap_or_else(|| {
let pooling_mode = config.get("pooling").and_then(|v| v.as_str());
pooling_type_u32(pooling_mode).unwrap_or_else(|| {
panic!(
"config.json key `pooling` has unrecognized value {pooling_mode:?}; \
expected one of mean | cls | last | none | rank"
)
})
});
let cls_out_labels: Option<Vec<String>> = config
.get("id2label")
.and_then(|v| v.as_object())
.and_then(|m| {
let mut entries: Vec<(i64, String)> = m
.iter()
.filter_map(|(k, v)| Some((k.parse::<i64>().ok()?, v.as_str()?.to_string())))
.collect();
if entries.is_empty() {
return None;
}
if entries.len() == 2 && entries.iter().any(|(k, v)| *k == 0 && v == "LABEL_0") {
return None;
}
entries.sort_by_key(|e| e.0);
Some(entries.into_iter().map(|(_, v)| v).collect())
});
let mut kv: Vec<(String, MetaValue)> = emit_general_prelude(
"bert",
display_name,
&id_components,
None,
model_card,
sampling,
);
kv.push(("bert.block_count".into(), MetaValue::U32(n_layers)));
kv.push(("bert.context_length".into(), MetaValue::U32(ctx_len)));
kv.push(("bert.embedding_length".into(), MetaValue::U32(hidden_size)));
kv.push(("bert.feed_forward_length".into(), MetaValue::U32(ffn_len)));
kv.push(("bert.attention.head_count".into(), MetaValue::U32(n_head)));
kv.push((
"bert.attention.layer_norm_epsilon".into(),
MetaValue::F32(ln_eps),
));
kv.push(("bert.attention.causal".into(), MetaValue::Bool(false)));
kv.push(("bert.pooling_type".into(), MetaValue::U32(pooling_u32)));
if let Some(labels) = cls_out_labels {
kv.push((
"bert.classifier.output_labels".into(),
MetaValue::ArrayString(labels),
));
}
kv.extend(emit_general_postlude(file_type));
kv
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn bert_tensor_name_round_trip() {
let cases: &[(&str, &str)] = &[
("embeddings.word_embeddings.weight", "token_embd.weight"),
(
"embeddings.position_embeddings.weight",
"position_embd.weight",
),
(
"embeddings.token_type_embeddings.weight",
"token_types.weight",
),
("embeddings.LayerNorm.weight", "token_embd_norm.weight"),
("embeddings.LayerNorm.bias", "token_embd_norm.bias"),
(
"encoder.layer.0.attention.self.query.weight",
"blk.0.attn_q.weight",
),
(
"encoder.layer.0.attention.self.query.bias",
"blk.0.attn_q.bias",
),
(
"encoder.layer.0.attention.self.key.weight",
"blk.0.attn_k.weight",
),
(
"encoder.layer.0.attention.self.key.bias",
"blk.0.attn_k.bias",
),
(
"encoder.layer.0.attention.self.value.weight",
"blk.0.attn_v.weight",
),
(
"encoder.layer.0.attention.self.value.bias",
"blk.0.attn_v.bias",
),
(
"encoder.layer.0.attention.output.dense.weight",
"blk.0.attn_output.weight",
),
(
"encoder.layer.0.attention.output.dense.bias",
"blk.0.attn_output.bias",
),
(
"encoder.layer.0.attention.output.LayerNorm.weight",
"blk.0.attn_output_norm.weight",
),
(
"encoder.layer.0.attention.output.LayerNorm.bias",
"blk.0.attn_output_norm.bias",
),
(
"encoder.layer.0.intermediate.dense.weight",
"blk.0.ffn_up.weight",
),
(
"encoder.layer.0.intermediate.dense.bias",
"blk.0.ffn_up.bias",
),
(
"encoder.layer.0.output.dense.weight",
"blk.0.ffn_down.weight",
),
("encoder.layer.0.output.dense.bias", "blk.0.ffn_down.bias"),
(
"encoder.layer.0.output.LayerNorm.weight",
"blk.0.layer_output_norm.weight",
),
(
"encoder.layer.0.output.LayerNorm.bias",
"blk.0.layer_output_norm.bias",
),
(
"encoder.layer.11.attention.self.query.weight",
"blk.11.attn_q.weight",
),
(
"encoder.layer.11.intermediate.dense.bias",
"blk.11.ffn_up.bias",
),
(
"encoder.layer.23.attention.output.LayerNorm.bias",
"blk.23.attn_output_norm.bias",
),
(
"encoder.layer.23.output.LayerNorm.weight",
"blk.23.layer_output_norm.weight",
),
("pooler.dense.weight", "cls.weight"),
("pooler.dense.bias", "cls.bias"),
];
for &(hf, expected_gguf) in cases {
let got = map_tensor_name(hf);
assert_eq!(
got.as_deref(),
Some(expected_gguf),
"map_tensor_name({hf:?}) = {got:?}, want Some({expected_gguf:?})"
);
}
}
#[test]
fn bert_tensor_name_strips_bert_prefix() {
let cases: &[(&str, &str)] = &[
(
"bert.embeddings.word_embeddings.weight",
"token_embd.weight",
),
("bert.embeddings.LayerNorm.bias", "token_embd_norm.bias"),
(
"bert.encoder.layer.5.attention.self.value.bias",
"blk.5.attn_v.bias",
),
(
"bert.encoder.layer.5.output.LayerNorm.weight",
"blk.5.layer_output_norm.weight",
),
("bert.pooler.dense.weight", "cls.weight"),
];
for &(hf, expected) in cases {
assert_eq!(
map_tensor_name(hf).as_deref(),
Some(expected),
"stripped-prefix mapping for {hf:?} failed"
);
}
}
#[test]
fn bert_tensor_name_rejects_unknown_kinds() {
assert_eq!(map_tensor_name("embeddings.unknown.weight"), None);
assert_eq!(map_tensor_name("transformer.h.0.attn.c_attn.weight"), None);
assert_eq!(
map_tensor_name("model.layers.0.self_attn.q_proj.weight"),
None
);
assert_eq!(
map_tensor_name("encoder.layer.01.attention.self.query.weight"),
None
);
assert_eq!(
map_tensor_name("encoder.layer..attention.self.query.weight"),
None
);
assert_eq!(
map_tensor_name("encoder.layer.attention.self.query.weight"),
None
);
assert_eq!(
map_tensor_name("encoder.layer.-1.attention.self.query.weight"),
None
);
assert_eq!(map_tensor_name("encoder.layer.0.unknown.weight"), None);
assert_eq!(
map_tensor_name("encoder.layer.0.attention.self.rotary_emb.inv_freq"),
None
);
assert_eq!(
map_tensor_name("encoder.layer.0.attention.self.query.gamma"),
None
);
}
#[test]
fn bert_metadata_built_from_config() {
let cfg = json!({
"_name_or_path": "BAAI/bge-large-en-v1.5",
"hidden_size": 1024,
"num_hidden_layers": 24,
"intermediate_size": 4096,
"num_attention_heads": 16,
"max_position_embeddings": 512,
"layer_norm_eps": 1.0e-12,
"pooling": "cls",
});
let kv = build_metadata(&cfg, 1 , None, 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("bert".into())
);
assert!(
matches!(by_key.get("general.name"), Some(MetaValue::String(_))),
"general.name must be present"
);
assert_eq!(by_key["bert.context_length"], MetaValue::U32(512));
assert_eq!(by_key["bert.embedding_length"], MetaValue::U32(1024));
assert_eq!(by_key["bert.block_count"], MetaValue::U32(24));
assert_eq!(by_key["bert.feed_forward_length"], MetaValue::U32(4096));
assert_eq!(by_key["bert.attention.head_count"], MetaValue::U32(16));
assert!(
by_key.get("bert.attention.head_count_kv").is_none(),
"canonical does NOT emit head_count_kv for BERT"
);
assert_eq!(
by_key["bert.attention.layer_norm_epsilon"],
MetaValue::F32(1.0e-12)
);
assert_eq!(
by_key["bert.attention.causal"],
MetaValue::Bool(false),
"BERT is encoder-only / bidirectional"
);
assert_eq!(
by_key["bert.pooling_type"],
MetaValue::U32(2),
"pooling=cls → PoolingType::CLS = 2"
);
assert_eq!(by_key["general.file_type"], MetaValue::U32(1));
assert_eq!(by_key["general.quantization_version"], MetaValue::U32(2));
}
#[test]
fn bert_metadata_optional_key_defaults() {
let cfg = json!({
"hidden_size": 768,
"num_hidden_layers": 12,
"intermediate_size": 3072,
"num_attention_heads": 12,
"max_position_embeddings": 512,
"layer_norm_eps": 1.0e-12,
});
let kv = build_metadata(&cfg, 0, None, 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()),
"name defaults to title-cased 'Model' when no source available"
);
assert_eq!(
by_key["bert.pooling_type"],
MetaValue::U32(1),
"pooling defaults to MEAN (=1)"
);
assert!(
by_key.get("bert.attention.head_count_kv").is_none(),
"canonical does NOT emit head_count_kv for BERT"
);
assert_eq!(
by_key["bert.attention.causal"],
MetaValue::Bool(false),
"causal=false even without explicit config opt-in"
);
}
}