use kopitiam_core::Result;
use kopitiam_loader::LoadedModel;
use kopitiam_tensor::Tensor;
use crate::bridge::{load_matmul_weight, load_matmul_weight_opt, load_tensor_f32, load_tensor_f32_opt};
use crate::config::QwenConfig;
pub(crate) struct LayerWeights {
pub attn_norm: Tensor,
pub wq: Tensor,
pub bq: Option<Tensor>,
pub wk: Tensor,
pub bk: Option<Tensor>,
pub wv: Tensor,
pub bv: Option<Tensor>,
pub wo: Tensor,
pub ffn_norm: Tensor,
pub w_gate: Tensor,
pub w_up: Tensor,
pub w_down: Tensor,
}
impl LayerWeights {
fn load(model: &LoadedModel, layer: usize) -> Result<Self> {
let p = |suffix: &str| format!("blk.{layer}.{suffix}");
Ok(Self {
attn_norm: load_tensor_f32(model, &p("attn_norm.weight"))?,
wq: load_matmul_weight(model, &p("attn_q.weight"))?,
bq: load_tensor_f32_opt(model, &p("attn_q.bias"))?,
wk: load_matmul_weight(model, &p("attn_k.weight"))?,
bk: load_tensor_f32_opt(model, &p("attn_k.bias"))?,
wv: load_matmul_weight(model, &p("attn_v.weight"))?,
bv: load_tensor_f32_opt(model, &p("attn_v.bias"))?,
wo: load_matmul_weight(model, &p("attn_output.weight"))?,
ffn_norm: load_tensor_f32(model, &p("ffn_norm.weight"))?,
w_gate: load_matmul_weight(model, &p("ffn_gate.weight"))?,
w_up: load_matmul_weight(model, &p("ffn_up.weight"))?,
w_down: load_matmul_weight(model, &p("ffn_down.weight"))?,
})
}
}
pub(crate) struct ModelWeights {
pub token_embd: Tensor,
pub layers: Vec<LayerWeights>,
pub output_norm: Tensor,
pub output_weight: Tensor,
}
impl ModelWeights {
pub(crate) fn load(model: &LoadedModel, config: &QwenConfig) -> Result<Self> {
let token_embd = load_tensor_f32(model, "token_embd.weight")?;
let layers = (0..config.n_layers).map(|i| LayerWeights::load(model, i)).collect::<Result<_>>()?;
let output_norm = load_tensor_f32(model, "output_norm.weight")?;
let output_weight = match load_matmul_weight_opt(model, "output.weight")? {
Some(w) => w,
None => token_embd.clone(),
};
Ok(Self { token_embd, layers, output_norm, output_weight })
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::synthetic_gguf::{build, write_temp_gguf, SyntheticModelSpec};
#[test]
fn loads_every_layer_and_ties_embeddings_when_output_weight_is_absent() {
let spec = SyntheticModelSpec { tie_embeddings: true, ..SyntheticModelSpec::default() };
let bytes = build(&spec);
let path = write_temp_gguf(&bytes, "weights-tied");
let model = kopitiam_loader::load_model(&path).unwrap();
let config = QwenConfig::from_metadata(model.metadata()).unwrap();
let weights = ModelWeights::load(&model, &config).unwrap();
assert_eq!(weights.layers.len(), spec.n_layers);
assert_eq!(weights.output_weight.shape(), weights.token_embd.shape());
assert_eq!(
weights.output_weight.to_vec_f32().unwrap(),
weights.token_embd.to_vec_f32().unwrap(),
"tied output weight must have the same values as the embedding table"
);
assert!(weights.layers[0].bq.is_some());
assert!(weights.layers[0].bk.is_some());
assert!(weights.layers[0].bv.is_some());
}
#[test]
fn loads_a_separate_output_weight_when_present() {
let spec = SyntheticModelSpec { tie_embeddings: false, ..SyntheticModelSpec::default() };
let bytes = build(&spec);
let path = write_temp_gguf(&bytes, "weights-untied");
let model = kopitiam_loader::load_model(&path).unwrap();
let config = QwenConfig::from_metadata(model.metadata()).unwrap();
let weights = ModelWeights::load(&model, &config).unwrap();
assert_ne!(weights.output_weight.to_vec_f32().unwrap(), weights.token_embd.to_vec_f32().unwrap());
}
#[test]
fn missing_qkv_bias_is_tolerated_when_the_spec_omits_it() {
let spec = SyntheticModelSpec { with_qkv_bias: false, ..SyntheticModelSpec::default() };
let bytes = build(&spec);
let path = write_temp_gguf(&bytes, "weights-no-bias");
let model = kopitiam_loader::load_model(&path).unwrap();
let config = QwenConfig::from_metadata(model.metadata()).unwrap();
let weights = ModelWeights::load(&model, &config).unwrap();
assert!(weights.layers[0].bq.is_none());
assert!(weights.layers[0].bk.is_none());
assert!(weights.layers[0].bv.is_none());
}
#[test]
fn matmul_operand_weights_stay_quantized_when_the_file_ships_them_quantized() {
let spec = SyntheticModelSpec {
quantize_matmul_weights: true,
tie_embeddings: false,
..SyntheticModelSpec::quantized_benchmark()
};
let bytes = build(&spec);
let path = write_temp_gguf(&bytes, "weights-quantized");
let model = kopitiam_loader::load_model(&path).unwrap();
let config = QwenConfig::from_metadata(model.metadata()).unwrap();
let weights = ModelWeights::load(&model, &config).unwrap();
assert_eq!(weights.layers[0].wq.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.layers[0].wk.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.layers[0].wv.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.layers[0].wo.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.layers[0].w_gate.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.layers[0].w_up.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.layers[0].w_down.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.output_weight.dtype(), kopitiam_core::DType::Q8_0);
assert_eq!(weights.token_embd.dtype(), kopitiam_core::DType::F32);
assert_eq!(weights.output_norm.dtype(), kopitiam_core::DType::F32);
assert_eq!(weights.layers[0].attn_norm.dtype(), kopitiam_core::DType::F32);
assert_eq!(weights.layers[0].bq.as_ref().unwrap().dtype(), kopitiam_core::DType::F32);
}
#[test]
fn matmul_operand_weights_stay_f32_when_the_file_ships_them_as_f32() {
let bytes = build(&SyntheticModelSpec::default());
let path = write_temp_gguf(&bytes, "weights-unquantized-native");
let model = kopitiam_loader::load_model(&path).unwrap();
let config = QwenConfig::from_metadata(model.metadata()).unwrap();
let weights = ModelWeights::load(&model, &config).unwrap();
assert_eq!(weights.layers[0].wq.dtype(), kopitiam_core::DType::F32);
}
}