use kopitiam_core::{Error, Result};
use kopitiam_loader::ModelMetadata;
const FORMAT: &str = "qwen-config";
#[derive(Debug, Clone, PartialEq)]
pub struct QwenConfig {
pub n_layers: usize,
pub n_heads: usize,
pub n_kv_heads: usize,
pub hidden_size: usize,
pub head_dim: usize,
pub ffn_hidden_size: usize,
pub vocab_size: usize,
pub max_context: usize,
pub rope_theta: f32,
pub rope_dimension_count: usize,
pub norm_eps: f32,
}
impl QwenConfig {
pub fn from_metadata(metadata: &ModelMetadata) -> Result<Self> {
let n_layers = require(metadata.n_layers, "block_count")?;
let n_heads = require(metadata.n_heads, "attention.head_count")? as usize;
let n_kv_heads = metadata.n_kv_heads.map(|v| v as usize).unwrap_or(n_heads);
let hidden_size = require(metadata.embedding_length, "embedding_length")? as usize;
let ffn_hidden_size = require(metadata.feed_forward_length, "feed_forward_length")?;
let vocab_size = require(metadata.vocab_size, "vocab_size (tokenizer.ggml.tokens length)")?;
let max_context = metadata.context_length.map(|v| v as usize).unwrap_or(4096);
if n_heads == 0 || !hidden_size.is_multiple_of(n_heads) {
return Err(malformed(format!(
"attention.head_count ({n_heads}) must evenly divide embedding_length ({hidden_size})"
)));
}
let head_dim = hidden_size / n_heads;
if n_kv_heads == 0 || !n_heads.is_multiple_of(n_kv_heads) {
return Err(malformed(format!(
"attention.head_count_kv ({n_kv_heads}) must evenly divide attention.head_count ({n_heads})"
)));
}
let rope_theta = metadata.rope_theta.unwrap_or(10_000.0);
let rope_dimension_count = metadata
.rope_dimension_count
.map(|v| v as usize)
.unwrap_or(head_dim);
if rope_dimension_count == 0 || rope_dimension_count > head_dim {
return Err(malformed(format!(
"rope.dimension_count ({rope_dimension_count}) must be in 1..=head_dim ({head_dim})"
)));
}
if !rope_dimension_count.is_multiple_of(2) {
return Err(malformed(format!(
"rope.dimension_count ({rope_dimension_count}) must be even: RoPE pairs dimension i with i + count/2"
)));
}
let norm_eps = metadata.norm_epsilon.unwrap_or(1e-6);
Ok(Self {
n_layers: n_layers as usize,
n_heads,
n_kv_heads,
hidden_size,
head_dim,
ffn_hidden_size: ffn_hidden_size as usize,
vocab_size: vocab_size as usize,
max_context,
rope_theta,
rope_dimension_count,
norm_eps,
})
}
pub fn gqa_group_size(&self) -> usize {
self.n_heads / self.n_kv_heads
}
}
fn require(field: Option<u64>, key: &str) -> Result<u64> {
field.ok_or_else(|| malformed(format!("missing required metadata field \"{key}\"")))
}
fn malformed(reason: String) -> Error {
Error::MalformedModel { format: FORMAT, reason }
}
#[cfg(test)]
mod tests {
use super::*;
fn base_metadata() -> ModelMetadata {
ModelMetadata {
n_layers: Some(2),
n_heads: Some(4),
n_kv_heads: Some(2),
embedding_length: Some(16),
feed_forward_length: Some(32),
context_length: Some(512),
vocab_size: Some(100),
..Default::default()
}
}
#[test]
fn resolves_a_complete_config_with_explicit_values() {
let mut metadata = base_metadata();
metadata.rope_theta = Some(1_000_000.0);
metadata.rope_dimension_count = Some(4);
metadata.norm_epsilon = Some(1e-5);
let config = QwenConfig::from_metadata(&metadata).unwrap();
assert_eq!(config.n_layers, 2);
assert_eq!(config.n_heads, 4);
assert_eq!(config.n_kv_heads, 2);
assert_eq!(config.hidden_size, 16);
assert_eq!(config.head_dim, 4);
assert_eq!(config.ffn_hidden_size, 32);
assert_eq!(config.vocab_size, 100);
assert_eq!(config.max_context, 512);
assert_eq!(config.rope_theta, 1_000_000.0);
assert_eq!(config.rope_dimension_count, 4);
assert_eq!(config.norm_eps, 1e-5);
assert_eq!(config.gqa_group_size(), 2);
}
#[test]
fn missing_n_kv_heads_falls_back_to_ordinary_multi_head_attention() {
let mut metadata = base_metadata();
metadata.n_kv_heads = None;
let config = QwenConfig::from_metadata(&metadata).unwrap();
assert_eq!(config.n_kv_heads, config.n_heads);
assert_eq!(config.gqa_group_size(), 1);
}
#[test]
fn missing_rope_and_norm_fields_use_documented_defaults() {
let metadata = base_metadata();
let config = QwenConfig::from_metadata(&metadata).unwrap();
assert_eq!(config.rope_theta, 10_000.0);
assert_eq!(config.rope_dimension_count, config.head_dim);
assert_eq!(config.norm_eps, 1e-6);
}
#[test]
fn missing_required_field_is_rejected() {
let mut metadata = base_metadata();
metadata.n_layers = None;
assert!(matches!(
QwenConfig::from_metadata(&metadata),
Err(Error::MalformedModel { .. })
));
}
#[test]
fn head_count_that_does_not_divide_hidden_size_is_rejected() {
let mut metadata = base_metadata();
metadata.n_heads = Some(3); assert!(matches!(
QwenConfig::from_metadata(&metadata),
Err(Error::MalformedModel { .. })
));
}
#[test]
fn kv_head_count_that_does_not_divide_head_count_is_rejected() {
let mut metadata = base_metadata();
metadata.n_kv_heads = Some(3); assert!(matches!(
QwenConfig::from_metadata(&metadata),
Err(Error::MalformedModel { .. })
));
}
#[test]
fn odd_rope_dimension_count_is_rejected() {
let mut metadata = base_metadata();
metadata.rope_dimension_count = Some(3);
assert!(matches!(
QwenConfig::from_metadata(&metadata),
Err(Error::MalformedModel { .. })
));
}
}