use crate::config::ModelConfig;
use crate::loader::{
load_f32_vec, load_f32_vec_optional, load_weight_matrix, slice_quantized_rows,
};
use crate::LoadError;
use ferrox_core::WeightMatrix;
use ferrox_gguf::{GgufError, TensorSource};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct FusedQkvRows {
q: usize,
kv: usize,
}
impl FusedQkvRows {
pub(crate) fn of(config: &ModelConfig) -> Self {
Self {
q: config.n_heads * config.head_dim,
kv: config.n_kv_heads * config.head_dim,
}
}
pub(crate) fn total(self) -> usize {
self.q + 2 * self.kv
}
pub(crate) fn spans(self) -> [(usize, usize); 3] {
[(0, self.q), (self.q, self.kv), (self.q + self.kv, self.kv)]
}
}
pub(crate) struct QkvProjections {
pub(crate) q: WeightMatrix,
pub(crate) k: WeightMatrix,
pub(crate) v: WeightMatrix,
pub(crate) q_bias: Option<Vec<f32>>,
pub(crate) k_bias: Option<Vec<f32>>,
pub(crate) v_bias: Option<Vec<f32>>,
}
pub(crate) fn load_fused_or_split_qkv(
file: &impl TensorSource,
layer: usize,
config: &ModelConfig,
) -> Result<QkvProjections, LoadError> {
let q_name = format!("blk.{layer}.attn_q.weight");
let k_name = format!("blk.{layer}.attn_k.weight");
let v_name = format!("blk.{layer}.attn_v.weight");
let fused_name = format!("blk.{layer}.attn_qkv.weight");
if file.find_tensor(&q_name).is_some() {
return Ok(QkvProjections {
q: load_weight_matrix(file, &q_name)?,
k: load_weight_matrix(file, &k_name)?,
v: load_weight_matrix(file, &v_name)?,
q_bias: load_f32_vec_optional(file, &format!("blk.{layer}.attn_q.bias"))?,
k_bias: load_f32_vec_optional(file, &format!("blk.{layer}.attn_k.bias"))?,
v_bias: load_f32_vec_optional(file, &format!("blk.{layer}.attn_v.bias"))?,
});
}
if file.find_tensor(&fused_name).is_none() {
return Err(LoadError::Gguf(GgufError::TensorNotFound(q_name)));
}
let rows = FusedQkvRows::of(config);
let fused = load_weight_matrix(file, &fused_name)?;
if fused.rows() != rows.total() {
return Err(LoadError::UnsupportedFeature(
config.name.to_string(),
format!(
"{fused_name} has {} rows; expected q+k+v = {} \
(n_heads*head_dim + 2*n_kv_heads*head_dim)",
fused.rows(),
rows.total()
),
));
}
let [q_span, k_span, v_span] = rows.spans();
let (q, k, v) = split_fused_weight(&fused, rows)?;
let (q_bias, k_bias, v_bias) = match split_fused_bias(file, layer, config, rows)? {
None => (None, None, None),
Some(b) => (
Some(b[q_span.0..q_span.0 + q_span.1].to_vec()),
Some(b[k_span.0..k_span.0 + k_span.1].to_vec()),
Some(b[v_span.0..v_span.0 + v_span.1].to_vec()),
),
};
Ok(QkvProjections {
q,
k,
v,
q_bias,
k_bias,
v_bias,
})
}
fn split_fused_weight(
fused: &WeightMatrix,
rows: FusedQkvRows,
) -> Result<(WeightMatrix, WeightMatrix, WeightMatrix), LoadError> {
let [q_span, k_span, v_span] = rows.spans();
if let (Some(q), Some(k), Some(v)) = (
slice_quantized_rows(fused, q_span.0, q_span.1),
slice_quantized_rows(fused, k_span.0, k_span.1),
slice_quantized_rows(fused, v_span.0, v_span.1),
) {
return Ok((q, k, v));
}
let cols = fused.cols();
let mut full = Vec::with_capacity(fused.rows() * cols);
for r in 0..fused.rows() {
full.extend_from_slice(&fused.dequant_row(r));
}
let take = |span: (usize, usize)| {
WeightMatrix::F32(ferrox_core::Tensor::new(
full[span.0 * cols..(span.0 + span.1) * cols].to_vec(),
vec![span.1, cols],
))
};
Ok((take(q_span), take(k_span), take(v_span)))
}
fn split_fused_bias(
file: &impl TensorSource,
layer: usize,
config: &ModelConfig,
rows: FusedQkvRows,
) -> Result<Option<Vec<f32>>, LoadError> {
let name = format!("blk.{layer}.attn_qkv.bias");
if file.find_tensor(&name).is_none() {
return Ok(None);
}
let bias = load_f32_vec(file, &name)?;
if bias.len() != rows.total() {
return Err(LoadError::UnsupportedFeature(
config.name.to_string(),
format!(
"{name} has {} elements; expected q+k+v = {} \
(n_heads*head_dim + 2*n_kv_heads*head_dim)",
bias.len(),
rows.total()
),
));
}
Ok(Some(bias))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_three_spans_tile_the_fused_tensor_exactly() {
for (n_heads, n_kv_heads, head_dim) in [(4, 2, 8), (32, 32, 128), (7, 1, 64)] {
let rows = FusedQkvRows {
q: n_heads * head_dim,
kv: n_kv_heads * head_dim,
};
let spans = rows.spans();
assert_eq!(spans[0].0, 0);
for w in spans.windows(2) {
assert_eq!(
w[0].0 + w[0].1,
w[1].0,
"spans must be contiguous for {n_heads}/{n_kv_heads}/{head_dim}"
);
}
let last = spans[2];
assert_eq!(
last.0 + last.1,
rows.total(),
"spans must end at total() for {n_heads}/{n_kv_heads}/{head_dim}"
);
}
}
#[test]
fn k_and_v_have_the_same_width_and_q_is_the_gqa_multiple() {
let rows = FusedQkvRows { q: 32, kv: 16 };
let spans = rows.spans();
assert_eq!(spans[1].1, spans[2].1);
assert_eq!(spans[0].1, 2 * spans[1].1);
assert_eq!(rows.total(), 64);
}
fn fused_qkv_gguf(
weight_rows: usize,
cols: usize,
bias_len: Option<usize>,
) -> ferrox_gguf::GgufFile {
let weight_bytes = weight_rows * cols * 4;
let mut plan = vec![ferrox_gguf::TensorPlan {
name: "blk.0.attn_qkv.weight".into(),
shape: vec![cols as u64, weight_rows as u64],
dtype: ferrox_gguf::GgmlType::F32,
byte_len: weight_bytes,
}];
if let Some(n) = bias_len {
plan.push(ferrox_gguf::TensorPlan {
name: "blk.0.attn_qkv.bias".into(),
shape: vec![n as u64],
dtype: ferrox_gguf::GgmlType::F32,
byte_len: n * 4,
});
}
let mut w = ferrox_gguf::GgufWriter::create(Vec::new(), &Default::default(), plan).unwrap();
w.write_tensor("blk.0.attn_qkv.weight", &vec![0u8; weight_bytes])
.unwrap();
if let Some(n) = bias_len {
w.write_tensor("blk.0.attn_qkv.bias", &vec![0u8; n * 4])
.unwrap();
}
let bytes = w.finish().unwrap();
let tag = bias_len.map_or("none".to_string(), |n| n.to_string());
let tmp =
std::env::temp_dir().join(format!("ferrox_qkv_fused_{weight_rows}_{cols}_{tag}.gguf"));
std::fs::write(&tmp, &bytes).unwrap();
let file = ferrox_gguf::GgufFile::open(&tmp).expect("the written file must parse");
std::fs::remove_file(&tmp).ok();
file
}
fn config_4x2_head8() -> ModelConfig {
let mut cfg = crate::config::glm_5_2();
cfg.n_heads = 4;
cfg.n_kv_heads = 2;
cfg.head_dim = 8;
cfg
}
#[test]
fn a_fused_bias_of_the_wrong_length_is_refused_by_name() {
let cfg = config_4x2_head8();
let rows = FusedQkvRows::of(&cfg);
assert_eq!(rows.total(), 64);
let file = fused_qkv_gguf(rows.total(), 8, Some(rows.total() - 1));
let Err(err) = load_fused_or_split_qkv(&file, 0, &cfg) else {
panic!("a fused bias one element short must refuse");
};
let msg = err.to_string();
assert!(msg.contains("attn_qkv.bias"), "{msg}");
assert!(
msg.contains("63"),
"the message must name what it found: {msg}"
);
assert!(msg.contains("64"), "and what it expected: {msg}");
}
#[test]
fn a_fused_bias_of_the_right_length_is_split_into_three() {
let cfg = config_4x2_head8();
let rows = FusedQkvRows::of(&cfg);
let file = fused_qkv_gguf(rows.total(), 8, Some(rows.total()));
let Ok(p) = load_fused_or_split_qkv(&file, 0, &cfg) else {
panic!("a fused bias of the declared length must load");
};
assert_eq!(p.q_bias.expect("q").len(), 32);
assert_eq!(p.k_bias.expect("k").len(), 16);
assert_eq!(p.v_bias.expect("v").len(), 16);
}
#[test]
fn a_fused_weight_with_no_bias_leaves_all_three_unset() {
let cfg = config_4x2_head8();
let rows = FusedQkvRows::of(&cfg);
let file = fused_qkv_gguf(rows.total(), 8, None);
assert!(file.find_tensor("blk.0.attn_qkv.bias").is_none());
let Ok(p) = load_fused_or_split_qkv(&file, 0, &cfg) else {
panic!("a fused weight with no bias must still load");
};
assert!(p.q_bias.is_none());
assert!(p.k_bias.is_none());
assert!(p.v_bias.is_none());
}
}