use crate::rope::RopeKind;
use kopitiam_core::{Error, Result};
use kopitiam_loader::ModelMetadata;
const SUPPORTED_ARCHITECTURES: &[&str] = &["llama", "qwen2"];
#[must_use]
pub fn is_supported_architecture(architecture: Option<&str>) -> bool {
architecture.is_some_and(|a| SUPPORTED_ARCHITECTURES.contains(&a))
}
#[must_use]
pub fn rope_kind_for_architecture(architecture: Option<&str>) -> RopeKind {
match architecture.unwrap_or_default() {
"llama" | "llama4" | "deci" | "baichuan" | "starcoder" | "internlm2" | "minicpm"
| "xverse" | "command-r" | "cohere2" | "olmo" | "arctic" | "deepseek" | "deepseek2" => {
RopeKind::Interleaved
}
"qwen2" | "qwen2moe" | "qwen3" | "qwen3moe" | "phi2" | "phi3" | "gemma" | "gemma2"
| "gemma3" | "stablelm" | "olmo2" | "gptneox" => RopeKind::SplitHalf,
_ => RopeKind::SplitHalf,
}
}
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 rope_kind: RopeKind,
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 architecture = metadata.architecture.as_deref();
if !is_supported_architecture(architecture) {
return Err(malformed(format!(
"unsupported general.architecture {:?}: this runtime implements the dense \
pre-norm RoPE/GQA/SwiGLU forward pass only, and executes {SUPPORTED_ARCHITECTURES:?}. \
Running other architectures on it produces confident nonsense rather than an \
error, so it is refused here instead",
architecture.unwrap_or("<absent>")
)));
}
let rope_kind = rope_kind_for_architecture(architecture);
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,
rope_kind,
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 {
architecture: Some("llama".to_string()),
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 { .. })
));
}
}
#[cfg(test)]
mod rope_kind_tests {
use super::*;
#[test]
fn llama_is_interleaved_and_qwen2_is_split_half() {
assert_eq!(rope_kind_for_architecture(Some("llama")), RopeKind::Interleaved);
assert_eq!(rope_kind_for_architecture(Some("qwen2")), RopeKind::SplitHalf);
}
#[test]
fn smollm2_reports_llama_and_therefore_gets_interleaved() {
assert_eq!(rope_kind_for_architecture(Some("llama")), RopeKind::Interleaved);
}
#[test]
fn an_unknown_architecture_falls_back_to_split_half() {
assert_eq!(rope_kind_for_architecture(None), RopeKind::SplitHalf);
assert_eq!(rope_kind_for_architecture(Some("not-a-real-arch")), RopeKind::SplitHalf);
}
}
#[cfg(test)]
mod architecture_gate_tests {
use super::*;
fn metadata_with(arch: Option<&str>) -> ModelMetadata {
ModelMetadata {
architecture: arch.map(str::to_string),
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 the_architectures_we_actually_execute_still_load() {
for arch in ["llama", "qwen2"] {
assert!(
QwenConfig::from_metadata(&metadata_with(Some(arch))).is_ok(),
"{arch} must still load"
);
}
}
#[test]
fn an_architecture_with_a_different_forward_pass_is_refused() {
for arch in ["gemma", "gemma3", "phi3", "qwen2moe", "mamba"] {
let err = QwenConfig::from_metadata(&metadata_with(Some(arch)))
.expect_err("{arch} must be refused");
assert!(matches!(err, Error::MalformedModel { .. }), "{arch}: {err:?}");
}
}
#[test]
fn an_absent_architecture_is_refused_rather_than_guessed() {
assert!(!is_supported_architecture(None));
assert!(QwenConfig::from_metadata(&metadata_with(None)).is_err());
}
#[test]
fn the_refusal_names_the_architecture_and_the_supported_set() {
let err = QwenConfig::from_metadata(&metadata_with(Some("gemma"))).unwrap_err();
let text = format!("{err}");
assert!(text.contains("gemma"), "should name the offender: {text}");
assert!(text.contains("llama") && text.contains("qwen2"), "should name what works: {text}");
}
}