use ferrox_core::tensor::Tensor;
use ferrox_core::weight_matrix::{QuantKind, WeightBytes, WeightMatrix};
use ferrox_gguf::{GgmlType, TensorSource};
use crate::glm_dsa::{Glm52AttnWeights, Glm52MlaConfig, IndexerConfig, IndexerWeights};
use crate::loader::LoadError;
fn find_info<'a>(
file: &'a impl TensorSource,
name: &str,
) -> Result<&'a ferrox_gguf::TensorInfo, LoadError> {
file.find_tensor(name)
.ok_or_else(|| LoadError::Gguf(ferrox_gguf::GgufError::TensorNotFound(name.to_string())))
}
fn load_f32_vec(file: &impl TensorSource, name: &str) -> Result<Vec<f32>, LoadError> {
let info = find_info(file, name)?;
let raw = file.tensor_bytes(name)?;
match info.dtype {
GgmlType::F32 => {
let mut out = Vec::with_capacity(raw.len() / 4);
for chunk in raw.chunks_exact(4) {
out.push(f32::from_le_bytes(chunk.try_into().unwrap()));
}
Ok(out)
}
GgmlType::F16 => ferrox_quant::dequant_f16(raw)
.map_err(|_| LoadError::UnsupportedDtype(name.to_string(), GgmlType::F16)),
GgmlType::BF16 => ferrox_quant::dequant_bf16(raw)
.map_err(|_| LoadError::UnsupportedDtype(name.to_string(), GgmlType::BF16)),
other => Err(LoadError::UnsupportedDtype(name.to_string(), other)),
}
}
fn load_weight_matrix(file: &impl TensorSource, name: &str) -> Result<WeightMatrix, LoadError> {
let info = find_info(file, name)?;
let shape: Vec<usize> = info.shape.iter().rev().map(|&d| d as usize).collect();
let (rows, cols) = match shape.as_slice() {
[r, c] => (*r, *c),
other => {
return Err(LoadError::UnsupportedDtype(
format!("{name} (expected 2D, got shape {other:?})"),
info.dtype,
))
}
};
match info.dtype {
GgmlType::F32 | GgmlType::F16 | GgmlType::BF16 => {
let data = load_f32_vec(file, name)?;
Ok(WeightMatrix::F32(Tensor::new(data, shape)))
}
other => match quant_kind_for(other) {
Some(kind) => {
let (mmap, range) = file.tensor_mapped_range(name)?;
Ok(WeightMatrix::Quantized {
data: WeightBytes::Mapped { mmap, range },
rows,
cols,
kind,
})
}
None => Err(LoadError::UnsupportedDtype(name.to_string(), other)),
},
}
}
fn quant_kind_for(dtype: GgmlType) -> Option<QuantKind> {
match dtype {
GgmlType::Q8_0 => Some(QuantKind::Q8_0),
GgmlType::Q4_0 => Some(QuantKind::Q4_0),
GgmlType::Q4K => Some(QuantKind::Q4K),
GgmlType::Q5K => Some(QuantKind::Q5K),
GgmlType::Q6K => Some(QuantKind::Q6K),
GgmlType::Q2K => Some(QuantKind::Q2K),
GgmlType::Q3K => Some(QuantKind::Q3K),
GgmlType::Q4_1 => Some(QuantKind::Q4_1),
GgmlType::Q5_0 => Some(QuantKind::Q5_0),
GgmlType::Q5_1 => Some(QuantKind::Q5_1),
GgmlType::Q8_1 => Some(QuantKind::Q8_1),
GgmlType::IQ4NL => Some(QuantKind::IQ4NL),
GgmlType::IQ4XS => Some(QuantKind::IQ4XS),
GgmlType::IQ2XS => Some(QuantKind::IQ2XS),
GgmlType::IQ2S => Some(QuantKind::IQ2S),
GgmlType::IQ3S => Some(QuantKind::IQ3S),
GgmlType::IQ1M => Some(QuantKind::IQ1M),
GgmlType::IQ1S => Some(QuantKind::IQ1S),
GgmlType::IQ2XXS => Some(QuantKind::IQ2XXS),
GgmlType::IQ3XXS => Some(QuantKind::IQ3XXS),
GgmlType::MXFP4 => Some(QuantKind::Mxfp4Gguf),
_ => None,
}
}
fn load_wk_b_transposed(
file: &impl TensorSource,
name: &str,
n_head: usize,
qk_nope_head_dim: usize,
kv_lora_rank: usize,
) -> Result<Vec<WeightMatrix>, LoadError> {
let info = find_info(file, name)?;
if info.shape.len() != 3
|| info.shape[0] as usize != qk_nope_head_dim
|| info.shape[1] as usize != kv_lora_rank
|| info.shape[2] as usize != n_head
{
return Err(LoadError::UnsupportedDtype(
format!(
"{name} (expected ne=[{qk_nope_head_dim}, {kv_lora_rank}, {n_head}], got {:?})",
info.shape
),
info.dtype,
));
}
if !matches!(info.dtype, GgmlType::F32 | GgmlType::F16 | GgmlType::BF16) {
return Err(LoadError::UnsupportedDtype(name.to_string(), info.dtype));
}
let all = load_f32_vec(file, name)?;
let per_head = kv_lora_rank * qk_nope_head_dim;
Ok((0..n_head)
.map(|h| {
let head_raw = &all[h * per_head..(h + 1) * per_head]; let mut transposed = vec![0f32; per_head]; for row in 0..kv_lora_rank {
for col in 0..qk_nope_head_dim {
transposed[col * kv_lora_rank + row] = head_raw[row * qk_nope_head_dim + col];
}
}
WeightMatrix::F32(Tensor::new(
transposed,
vec![qk_nope_head_dim, kv_lora_rank],
))
})
.collect())
}
fn load_wv_b(
file: &impl TensorSource,
name: &str,
n_head: usize,
kv_lora_rank: usize,
v_head_dim: usize,
) -> Result<Vec<WeightMatrix>, LoadError> {
let info = find_info(file, name)?;
if info.shape.len() != 3
|| info.shape[0] as usize != kv_lora_rank
|| info.shape[1] as usize != v_head_dim
|| info.shape[2] as usize != n_head
{
return Err(LoadError::UnsupportedDtype(
format!(
"{name} (expected ne=[{kv_lora_rank}, {v_head_dim}, {n_head}], got {:?})",
info.shape
),
info.dtype,
));
}
match info.dtype {
GgmlType::F32 | GgmlType::F16 | GgmlType::BF16 => {
let all = load_f32_vec(file, name)?;
let per_head = kv_lora_rank * v_head_dim;
Ok((0..n_head)
.map(|h| {
WeightMatrix::F32(Tensor::new(
all[h * per_head..(h + 1) * per_head].to_vec(),
vec![v_head_dim, kv_lora_rank],
))
})
.collect())
}
other => match quant_kind_for(other) {
Some(kind) => {
let (mmap, full_range) = file.tensor_mapped_range(name)?;
let bytes_per_head = (full_range.end - full_range.start) / n_head;
Ok((0..n_head)
.map(|h| WeightMatrix::Quantized {
data: WeightBytes::Mapped {
mmap: std::sync::Arc::clone(&mmap),
range: (full_range.start + h * bytes_per_head)
..(full_range.start + (h + 1) * bytes_per_head),
},
rows: v_head_dim,
cols: kv_lora_rank,
kind,
})
.collect())
}
None => Err(LoadError::UnsupportedDtype(name.to_string(), other)),
},
}
}
fn split_expert_tensor(
file: &impl TensorSource,
name: &str,
n_experts: usize,
) -> Result<Vec<WeightMatrix>, LoadError> {
let info = find_info(file, name)?;
if info.shape.len() != 3 || info.shape[2] as usize != n_experts {
let file_experts = info.shape.last().map(|&d| d as usize).unwrap_or(0);
return Err(LoadError::ExpertCountMismatch(
name.to_string(),
file_experts,
n_experts,
));
}
let out_dim = info.shape[1] as usize;
let in_dim = info.shape[0] as usize;
let raw = file.tensor_bytes(name)?;
match info.dtype {
GgmlType::F32 | GgmlType::F16 | GgmlType::BF16 => {
let all = crate::loader::widen_plain_float(info.dtype, raw, name)?;
let per_expert = out_dim * in_dim;
Ok((0..n_experts)
.map(|e| {
WeightMatrix::F32(Tensor::new(
all[e * per_expert..(e + 1) * per_expert].to_vec(),
vec![out_dim, in_dim],
))
})
.collect())
}
other => match quant_kind_for(other) {
Some(kind) => {
let (mmap, full_range) = file.tensor_mapped_range(name)?;
let bytes_per_expert = raw.len() / n_experts;
Ok((0..n_experts)
.map(|e| WeightMatrix::Quantized {
data: WeightBytes::Mapped {
mmap: std::sync::Arc::clone(&mmap),
range: (full_range.start + e * bytes_per_expert)
..(full_range.start + (e + 1) * bytes_per_expert),
},
rows: out_dim,
cols: in_dim,
kind,
})
.collect())
}
None => Err(LoadError::UnsupportedDtype(name.to_string(), other)),
},
}
}
pub struct Glm52GgufHparams {
pub hidden_dim: usize,
pub num_heads: usize,
pub q_lora_rank: usize,
pub kv_lora_rank: usize,
pub qk_nope_head_dim: usize,
pub qk_rope_head_dim: usize,
pub v_head_dim: usize,
pub rope_theta: f32,
pub indexer_n_heads: usize,
pub indexer_head_dim: usize,
pub indexer_rope_dim: usize,
pub indexer_top_k: usize,
pub dense_ffn_dim: usize,
pub moe_ffn_dim: usize,
pub n_experts: usize,
pub n_shared_experts: usize,
}
pub fn load_glm52_attn(
file: &impl TensorSource,
hp: &Glm52GgufHparams,
layer_idx: usize,
is_full_indexer_layer: bool,
) -> Result<Glm52AttnWeights, LoadError> {
let l = layer_idx;
let q_head_dim = hp.qk_nope_head_dim + hp.qk_rope_head_dim;
let q_a_proj = load_weight_matrix(file, &format!("blk.{l}.attn_q_a.weight"))?;
assert_eq!(
q_a_proj.rows(),
hp.q_lora_rank,
"blk.{l}.attn_q_a.weight row count"
);
assert_eq!(
q_a_proj.cols(),
hp.hidden_dim,
"blk.{l}.attn_q_a.weight col count"
);
let q_b_proj = load_weight_matrix(file, &format!("blk.{l}.attn_q_b.weight"))?;
assert_eq!(
q_b_proj.rows(),
hp.num_heads * q_head_dim,
"blk.{l}.attn_q_b.weight row count"
);
let kv_a_proj_with_mqa = load_weight_matrix(file, &format!("blk.{l}.attn_kv_a_mqa.weight"))?;
assert_eq!(
kv_a_proj_with_mqa.rows(),
hp.kv_lora_rank + hp.qk_rope_head_dim,
"blk.{l}.attn_kv_a_mqa.weight row count"
);
let wk_b = load_wk_b_transposed(
file,
&format!("blk.{l}.attn_k_b.weight"),
hp.num_heads,
hp.qk_nope_head_dim,
hp.kv_lora_rank,
)?;
let wv_b = load_wv_b(
file,
&format!("blk.{l}.attn_v_b.weight"),
hp.num_heads,
hp.kv_lora_rank,
hp.v_head_dim,
)?;
let o_proj = load_weight_matrix(file, &format!("blk.{l}.attn_output.weight"))?;
assert_eq!(
o_proj.rows(),
hp.hidden_dim,
"blk.{l}.attn_output.weight row count"
);
assert_eq!(
o_proj.cols(),
hp.num_heads * hp.v_head_dim,
"blk.{l}.attn_output.weight col count"
);
let indexer = if is_full_indexer_layer {
let k_norm_weight = load_f32_vec(file, &format!("blk.{l}.indexer.k_norm.weight"))?;
let k_norm_bias = load_f32_vec(file, &format!("blk.{l}.indexer.k_norm.bias"))?;
let proj = load_weight_matrix(file, &format!("blk.{l}.indexer.proj.weight"))?;
assert_eq!(
proj.rows(),
hp.indexer_n_heads,
"blk.{l}.indexer.proj.weight row count"
);
let attn_k = load_weight_matrix(file, &format!("blk.{l}.indexer.attn_k.weight"))?;
assert_eq!(
attn_k.rows(),
hp.indexer_head_dim,
"blk.{l}.indexer.attn_k.weight row count"
);
let attn_q_b = load_weight_matrix(file, &format!("blk.{l}.indexer.attn_q_b.weight"))?;
assert_eq!(
attn_q_b.rows(),
hp.indexer_n_heads * hp.indexer_head_dim,
"blk.{l}.indexer.attn_q_b.weight row count"
);
Some(IndexerWeights {
k_norm_weight,
k_norm_bias,
proj,
attn_k,
attn_q_b,
})
} else {
None
};
Ok(Glm52AttnWeights {
q_a_proj,
q_a_layernorm: load_f32_vec(file, &format!("blk.{l}.attn_q_a_norm.weight"))?,
q_b_proj,
kv_a_proj_with_mqa,
kv_a_layernorm: load_f32_vec(file, &format!("blk.{l}.attn_kv_a_norm.weight"))?,
wk_b,
wv_b,
o_proj,
indexer,
})
}
pub fn glm52_mla_config(hp: &Glm52GgufHparams) -> Glm52MlaConfig {
Glm52MlaConfig {
num_heads: hp.num_heads,
q_lora_rank: hp.q_lora_rank,
kv_lora_rank: hp.kv_lora_rank,
qk_nope_head_dim: hp.qk_nope_head_dim,
qk_rope_head_dim: hp.qk_rope_head_dim,
v_head_dim: hp.v_head_dim,
rope: crate::config::MlaRopeConfig {
theta: hp.rope_theta,
},
}
}
pub fn glm52_indexer_config(hp: &Glm52GgufHparams) -> IndexerConfig {
IndexerConfig {
n_heads: hp.indexer_n_heads,
head_dim: hp.indexer_head_dim,
rope_dim: hp.indexer_rope_dim,
top_k: hp.indexer_top_k,
rope_theta: hp.rope_theta,
}
}
pub struct Glm52DenseFfnWeights {
pub gate_proj: WeightMatrix,
pub up_proj: WeightMatrix,
pub down_proj: WeightMatrix,
}
pub fn load_glm52_dense_ffn(
file: &impl TensorSource,
layer_idx: usize,
) -> Result<Glm52DenseFfnWeights, LoadError> {
let l = layer_idx;
Ok(Glm52DenseFfnWeights {
gate_proj: load_weight_matrix(file, &format!("blk.{l}.ffn_gate.weight"))?,
up_proj: load_weight_matrix(file, &format!("blk.{l}.ffn_up.weight"))?,
down_proj: load_weight_matrix(file, &format!("blk.{l}.ffn_down.weight"))?,
})
}
pub struct Glm52MoeFfnWeights {
pub router_weight: WeightMatrix,
pub e_score_correction_bias: Vec<f32>,
pub experts: Vec<ferrox_moe::ExpertWeights>,
pub shared_expert: ferrox_moe::ExpertWeights,
}
pub fn load_glm52_moe_ffn(
file: &impl TensorSource,
hp: &Glm52GgufHparams,
layer_idx: usize,
) -> Result<Glm52MoeFfnWeights, LoadError> {
let l = layer_idx;
let gate_exps =
split_expert_tensor(file, &format!("blk.{l}.ffn_gate_exps.weight"), hp.n_experts)?;
let down_exps =
split_expert_tensor(file, &format!("blk.{l}.ffn_down_exps.weight"), hp.n_experts)?;
let up_exps = split_expert_tensor(file, &format!("blk.{l}.ffn_up_exps.weight"), hp.n_experts)?;
let experts = gate_exps
.into_iter()
.zip(down_exps)
.zip(up_exps)
.map(|((gate, down), up)| ferrox_moe::ExpertWeights { gate, up, down })
.collect();
let shared_expert = ferrox_moe::ExpertWeights {
gate: load_weight_matrix(file, &format!("blk.{l}.ffn_gate_shexp.weight"))?,
up: load_weight_matrix(file, &format!("blk.{l}.ffn_up_shexp.weight"))?,
down: load_weight_matrix(file, &format!("blk.{l}.ffn_down_shexp.weight"))?,
};
Ok(Glm52MoeFfnWeights {
router_weight: load_weight_matrix(file, &format!("blk.{l}.ffn_gate_inp.weight"))?,
e_score_correction_bias: load_f32_vec(file, &format!("blk.{l}.ffn_exp_probs_b.bias"))?,
experts,
shared_expert,
})
}
fn meta_u64(file: &impl TensorSource, key: &str) -> Result<u64, LoadError> {
file.metadata_u64(key)
.ok_or_else(|| LoadError::MissingHparam(key.to_string()))
}
fn meta_f32(file: &impl TensorSource, key: &str, default: f32) -> f32 {
file.metadata_f32(key).unwrap_or(default)
}
pub struct Glm52FileMeta {
pub arch: String,
pub n_layer: usize,
pub leading_dense: usize,
pub rms_norm_eps: f32,
pub n_experts_active: usize,
pub moe_renormalize: bool,
pub routed_scaling_factor: f32,
}
pub fn read_glm52_hparams(
file: &impl TensorSource,
) -> Result<(Glm52GgufHparams, Glm52FileMeta), LoadError> {
let arch = file
.metadata_str("general.architecture")
.ok_or_else(|| LoadError::MissingHparam("general.architecture".into()))?
.to_string();
if !matches!(arch.as_str(), "glm-dsa" | "glm4" | "glm4moe") {
return Err(LoadError::UnsupportedArchitecture(arch));
}
let p = |suffix: &str| format!("{arch}.{suffix}");
let n_layer = meta_u64(file, &p("block_count"))? as usize;
let hidden_dim = meta_u64(file, &p("embedding_length"))? as usize;
let dense_ffn_dim = meta_u64(file, &p("feed_forward_length"))? as usize;
let moe_ffn_dim = file
.metadata_u64(&p("expert_feed_forward_length"))
.unwrap_or(dense_ffn_dim as u64) as usize;
let n_heads = meta_u64(file, &p("attention.head_count"))? as usize;
let q_lora_rank = meta_u64(file, &p("attention.q_lora_rank"))? as usize;
let kv_lora_rank = meta_u64(file, &p("attention.kv_lora_rank"))? as usize;
let qk_nope_head_dim = meta_u64(file, &p("attention.qk_nope_head_dim"))? as usize;
let qk_rope_head_dim = meta_u64(file, &p("attention.qk_rope_head_dim"))? as usize;
let v_head_dim = file
.metadata_u64(&p("attention.v_head_dim"))
.or_else(|| file.metadata_u64(&p("attention.key_length")))
.unwrap_or(qk_nope_head_dim as u64) as usize;
let leading_dense = file
.metadata_u64(&p("leading_dense_block_count"))
.unwrap_or(0) as usize;
let n_experts = file.metadata_u64(&p("expert_count")).unwrap_or(0) as usize;
let n_shared_experts = file.metadata_u64(&p("expert_shared_count")).unwrap_or(1) as usize;
let n_experts_active = file.metadata_u64(&p("expert_used_count")).unwrap_or(8) as usize;
let indexer_n_heads = file
.metadata_u64(&p("attention.indexer_n_heads"))
.unwrap_or(4) as usize;
let indexer_head_dim = file
.metadata_u64(&p("attention.indexer_head_dim"))
.unwrap_or(128) as usize;
let indexer_top_k = file
.metadata_u64(&p("attention.indexer_top_k"))
.unwrap_or(2048) as usize;
let hp = Glm52GgufHparams {
hidden_dim,
num_heads: n_heads,
q_lora_rank,
kv_lora_rank,
qk_nope_head_dim,
qk_rope_head_dim,
v_head_dim,
rope_theta: meta_f32(file, &p("rope.freq_base"), 1_000_000.0),
indexer_n_heads,
indexer_head_dim,
indexer_rope_dim: qk_rope_head_dim,
indexer_top_k,
dense_ffn_dim,
moe_ffn_dim,
n_experts,
n_shared_experts,
};
let meta = Glm52FileMeta {
arch: arch.clone(),
n_layer,
leading_dense: leading_dense.min(n_layer),
rms_norm_eps: meta_f32(file, &p("attention.layer_norm_rms_epsilon"), 1e-5),
n_experts_active,
moe_renormalize: file
.metadata_u64(&p("expert_norm_topk_prob"))
.is_some_and(|v| v != 0),
routed_scaling_factor: meta_f32(file, &p("expert_routing_scale"), 2.5),
};
Ok((hp, meta))
}
fn is_full_indexer_layer(file: &impl TensorSource, layer_idx: usize) -> bool {
file.find_tensor(&format!("blk.{layer_idx}.indexer.proj.weight"))
.is_some()
}
fn is_dense_ffn_layer(file: &impl TensorSource, layer_idx: usize, leading_dense: usize) -> bool {
if layer_idx < leading_dense {
return true;
}
file.find_tensor(&format!("blk.{layer_idx}.ffn_gate.weight"))
.is_some()
&& file
.find_tensor(&format!("blk.{layer_idx}.ffn_gate_inp.weight"))
.is_none()
}
fn load_embedding_tensor(
file: &impl TensorSource,
hidden_dim: usize,
) -> Result<ferrox_core::tensor::Tensor, LoadError> {
let wm = load_weight_matrix(file, "token_embd.weight")?;
let vocab = wm.rows();
assert_eq!(wm.cols(), hidden_dim, "token_embd.weight col count");
let mut data = vec![0f32; vocab * hidden_dim];
for row in 0..vocab {
let r = wm.dequant_row(row);
data[row * hidden_dim..(row + 1) * hidden_dim].copy_from_slice(&r);
}
Ok(ferrox_core::tensor::Tensor::new(
data,
vec![vocab, hidden_dim],
))
}
pub fn load_glm52_engine(
file: &impl TensorSource,
) -> Result<crate::engine::Glm52Engine, LoadError> {
use crate::glm52_decoder::{
Glm52DecoderConfig, Glm52DecoderLayerWeights, Glm52DecoderWeights,
Glm52DenseFfnWeights as DecDenseFfn, Glm52LayerFfn, Glm52MoeFfnWeights as DecMoeFfn,
};
let (hp, meta) = read_glm52_hparams(file)?;
let embedding = load_embedding_tensor(file, hp.hidden_dim)?;
let final_norm_weight = load_f32_vec(file, "output_norm.weight")?;
let output_head = match load_weight_matrix(file, "output.weight") {
Ok(w) => w,
Err(_) => load_weight_matrix(file, "token_embd.weight")?,
};
let mut layers = Vec::with_capacity(meta.n_layer);
for layer_idx in 0..meta.n_layer {
let is_full = is_full_indexer_layer(file, layer_idx);
let attn = load_glm52_attn(file, &hp, layer_idx, is_full)?;
let ffn = if is_dense_ffn_layer(file, layer_idx, meta.leading_dense) {
let d = load_glm52_dense_ffn(file, layer_idx)?;
Glm52LayerFfn::Dense(Box::new(DecDenseFfn {
gate_proj: d.gate_proj,
up_proj: d.up_proj,
down_proj: d.down_proj,
}))
} else {
let m = load_glm52_moe_ffn(file, &hp, layer_idx)?;
Glm52LayerFfn::Moe(Box::new(DecMoeFfn {
router_weight: m.router_weight,
e_score_correction_bias: m.e_score_correction_bias,
experts: m.experts,
shared_expert: m.shared_expert,
}))
};
layers.push(Glm52DecoderLayerWeights {
attn_norm_weight: load_f32_vec(file, &format!("blk.{layer_idx}.attn_norm.weight"))?,
attn,
ffn_norm_weight: load_f32_vec(file, &format!("blk.{layer_idx}.ffn_norm.weight"))?,
ffn,
is_full_indexer_layer: is_full,
});
}
let weights = Glm52DecoderWeights {
embedding,
layers,
final_norm_weight,
output_head,
};
let cfg = Glm52DecoderConfig {
rms_norm_eps: meta.rms_norm_eps,
mla: glm52_mla_config(&hp),
indexer: glm52_indexer_config(&hp),
n_experts_active: meta.n_experts_active,
moe_renormalize: meta.moe_renormalize,
routed_scaling_factor: meta.routed_scaling_factor,
};
Ok(crate::engine::Glm52Engine { weights, cfg })
}
#[cfg(test)]
mod tests {
use super::*;
use byteorder::{LittleEndian, WriteBytesExt};
use std::io::Write;
fn f32_bytes(values: &[f32]) -> Vec<u8> {
values.iter().flat_map(|v| v.to_le_bytes()).collect()
}
struct FixtureTensor {
name: String,
shape: Vec<u64>,
bytes: Vec<u8>,
}
fn f32_tensor(name: impl Into<String>, shape: Vec<u64>, values: Vec<f32>) -> FixtureTensor {
FixtureTensor {
name: name.into(),
shape,
bytes: f32_bytes(&values),
}
}
fn build_gguf(arch: &str, tensors: &[FixtureTensor]) -> Vec<u8> {
let mut buf = Vec::new();
buf.write_u32::<LittleEndian>(ferrox_gguf::GGUF_MAGIC)
.unwrap();
buf.write_u32::<LittleEndian>(3).unwrap(); buf.write_u64::<LittleEndian>(tensors.len() as u64).unwrap();
buf.write_u64::<LittleEndian>(1).unwrap();
let write_string = |buf: &mut Vec<u8>, s: &str| {
buf.write_u64::<LittleEndian>(s.len() as u64).unwrap();
buf.write_all(s.as_bytes()).unwrap();
};
write_string(&mut buf, "general.architecture");
buf.write_u32::<LittleEndian>(8).unwrap(); write_string(&mut buf, arch);
let mut offset = 0u64;
let mut offsets = Vec::with_capacity(tensors.len());
for t in tensors {
write_string(&mut buf, &t.name);
buf.write_u32::<LittleEndian>(t.shape.len() as u32).unwrap();
for &d in t.shape.iter().rev() {
buf.write_u64::<LittleEndian>(d).unwrap();
}
buf.write_u32::<LittleEndian>(0).unwrap(); offsets.push(offset);
buf.write_u64::<LittleEndian>(offset).unwrap();
let padded = t.bytes.len().div_ceil(32) * 32;
offset += padded as u64;
}
while buf.len() % 32 != 0 {
buf.push(0);
}
let data_start = buf.len();
for (t, &off) in tensors.iter().zip(offsets.iter()) {
let want_len = data_start + off as usize;
while buf.len() < want_len {
buf.push(0);
}
buf.extend_from_slice(&t.bytes);
while buf.len() % 32 != 0 {
buf.push(0);
}
}
buf
}
struct Dims {
hidden_dim: usize,
num_heads: usize,
q_lora_rank: usize,
kv_lora_rank: usize,
qk_nope_head_dim: usize,
qk_rope_head_dim: usize,
v_head_dim: usize,
indexer_n_heads: usize,
indexer_head_dim: usize,
dense_ffn_dim: usize,
moe_ffn_dim: usize,
n_experts: usize,
n_shared_experts: usize,
}
fn push_layer_tensors(
tensors: &mut Vec<FixtureTensor>,
l: usize,
is_full: bool,
is_dense: bool,
d: &Dims,
) {
let h = d.hidden_dim;
let q_head_dim = d.qk_nope_head_dim + d.qk_rope_head_dim;
tensors.push(f32_tensor(
format!("blk.{l}.attn_norm.weight"),
vec![h as u64],
vec![1.0; h],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_norm.weight"),
vec![h as u64],
vec![1.0; h],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_q_a_norm.weight"),
vec![d.q_lora_rank as u64],
vec![1.0; d.q_lora_rank],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_kv_a_norm.weight"),
vec![d.kv_lora_rank as u64],
vec![1.0; d.kv_lora_rank],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_q_a.weight"),
vec![d.q_lora_rank as u64, h as u64],
vec![0.02; d.q_lora_rank * h],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_q_b.weight"),
vec![(d.num_heads * q_head_dim) as u64, d.q_lora_rank as u64],
vec![0.02; d.num_heads * q_head_dim * d.q_lora_rank],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_kv_a_mqa.weight"),
vec![(d.kv_lora_rank + d.qk_rope_head_dim) as u64, h as u64],
vec![0.02; (d.kv_lora_rank + d.qk_rope_head_dim) * h],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_k_b.weight"),
vec![
d.num_heads as u64,
d.kv_lora_rank as u64,
d.qk_nope_head_dim as u64,
],
vec![0.02; d.num_heads * d.kv_lora_rank * d.qk_nope_head_dim],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_v_b.weight"),
vec![
d.num_heads as u64,
d.v_head_dim as u64,
d.kv_lora_rank as u64,
],
vec![0.02; d.num_heads * d.v_head_dim * d.kv_lora_rank],
));
tensors.push(f32_tensor(
format!("blk.{l}.attn_output.weight"),
vec![h as u64, (d.num_heads * d.v_head_dim) as u64],
vec![0.02; h * d.num_heads * d.v_head_dim],
));
if is_full {
tensors.push(f32_tensor(
format!("blk.{l}.indexer.k_norm.weight"),
vec![d.indexer_head_dim as u64],
vec![1.0; d.indexer_head_dim],
));
tensors.push(f32_tensor(
format!("blk.{l}.indexer.k_norm.bias"),
vec![d.indexer_head_dim as u64],
vec![0.0; d.indexer_head_dim],
));
tensors.push(f32_tensor(
format!("blk.{l}.indexer.proj.weight"),
vec![d.indexer_n_heads as u64, h as u64],
vec![0.02; d.indexer_n_heads * h],
));
tensors.push(f32_tensor(
format!("blk.{l}.indexer.attn_k.weight"),
vec![d.indexer_head_dim as u64, h as u64],
vec![0.02; d.indexer_head_dim * h],
));
tensors.push(f32_tensor(
format!("blk.{l}.indexer.attn_q_b.weight"),
vec![
(d.indexer_n_heads * d.indexer_head_dim) as u64,
d.q_lora_rank as u64,
],
vec![0.02; d.indexer_n_heads * d.indexer_head_dim * d.q_lora_rank],
));
}
if is_dense {
for name in ["ffn_gate", "ffn_up"] {
tensors.push(f32_tensor(
format!("blk.{l}.{name}.weight"),
vec![d.dense_ffn_dim as u64, h as u64],
vec![0.02; d.dense_ffn_dim * h],
));
}
tensors.push(f32_tensor(
format!("blk.{l}.ffn_down.weight"),
vec![h as u64, d.dense_ffn_dim as u64],
vec![0.02; h * d.dense_ffn_dim],
));
} else {
let ff = d.moe_ffn_dim;
let n = d.n_experts;
tensors.push(f32_tensor(
format!("blk.{l}.ffn_gate_inp.weight"),
vec![n as u64, h as u64],
vec![0.02; n * h],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_exp_probs_b.bias"),
vec![n as u64],
vec![0.0; n],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_gate_exps.weight"),
vec![n as u64, ff as u64, h as u64],
vec![0.02; h * ff * n],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_down_exps.weight"),
vec![n as u64, h as u64, ff as u64],
vec![0.02; ff * h * n],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_up_exps.weight"),
vec![n as u64, ff as u64, h as u64],
vec![0.02; h * ff * n],
));
let shexp_dim = ff * d.n_shared_experts;
tensors.push(f32_tensor(
format!("blk.{l}.ffn_gate_shexp.weight"),
vec![shexp_dim as u64, h as u64],
vec![0.02; shexp_dim * h],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_down_shexp.weight"),
vec![h as u64, shexp_dim as u64],
vec![0.02; h * shexp_dim],
));
tensors.push(f32_tensor(
format!("blk.{l}.ffn_up_shexp.weight"),
vec![shexp_dim as u64, h as u64],
vec![0.02; shexp_dim * h],
));
}
}
fn dims() -> Dims {
Dims {
hidden_dim: 8,
num_heads: 2,
q_lora_rank: 6,
kv_lora_rank: 4,
qk_nope_head_dim: 4,
qk_rope_head_dim: 4,
v_head_dim: 3,
indexer_n_heads: 2,
indexer_head_dim: 4,
dense_ffn_dim: 5,
moe_ffn_dim: 4,
n_experts: 3,
n_shared_experts: 1,
}
}
fn hp_from(d: &Dims) -> Glm52GgufHparams {
Glm52GgufHparams {
hidden_dim: d.hidden_dim,
num_heads: d.num_heads,
q_lora_rank: d.q_lora_rank,
kv_lora_rank: d.kv_lora_rank,
qk_nope_head_dim: d.qk_nope_head_dim,
qk_rope_head_dim: d.qk_rope_head_dim,
v_head_dim: d.v_head_dim,
rope_theta: 8_000_000.0,
indexer_n_heads: d.indexer_n_heads,
indexer_head_dim: d.indexer_head_dim,
indexer_rope_dim: d.qk_rope_head_dim,
indexer_top_k: 2,
dense_ffn_dim: d.dense_ffn_dim,
moe_ffn_dim: d.moe_ffn_dim,
n_experts: d.n_experts,
n_shared_experts: d.n_shared_experts,
}
}
#[test]
fn loads_a_full_indexer_dense_layer_and_a_shared_indexer_moe_layer() {
let d = dims();
let mut tensors: Vec<FixtureTensor> = Vec::new();
push_layer_tensors(&mut tensors, 0, true, true, &d);
push_layer_tensors(&mut tensors, 1, false, false, &d);
let bytes = build_gguf("glm-dsa", &tensors);
let path = std::env::temp_dir().join(format!(
"ferrox_glm52_gguf_test_{}.gguf",
std::process::id()
));
std::fs::write(&path, &bytes).unwrap();
let file = ferrox_gguf::GgufFile::open(&path).expect("synthetic GGUF must parse");
let hp = hp_from(&d);
let layer0 = load_glm52_attn(&file, &hp, 0, true).expect("full-indexer layer must load");
assert!(layer0.indexer.is_some());
assert_eq!(layer0.wk_b.len(), d.num_heads);
assert_eq!(layer0.wk_b[0].rows(), d.qk_nope_head_dim);
assert_eq!(layer0.wk_b[0].cols(), d.kv_lora_rank);
assert_eq!(layer0.wv_b[0].rows(), d.v_head_dim);
assert_eq!(layer0.wv_b[0].cols(), d.kv_lora_rank);
let dense0 = load_glm52_dense_ffn(&file, 0).expect("dense FFN must load");
assert_eq!(dense0.gate_proj.rows(), d.dense_ffn_dim);
let layer1 = load_glm52_attn(&file, &hp, 1, false).expect("shared-indexer layer must load");
assert!(
layer1.indexer.is_none(),
"a \"shared\" layer must not load its own indexer weights"
);
let moe1 = load_glm52_moe_ffn(&file, &hp, 1).expect("MoE FFN must load");
assert_eq!(moe1.experts.len(), d.n_experts);
assert_eq!(moe1.e_score_correction_bias.len(), d.n_experts);
std::fs::remove_file(&path).ok();
}
#[test]
fn wk_b_transpose_matches_hand_computed_values() {
let tensors = vec![f32_tensor(
"blk.0.attn_k_b.weight",
vec![1, 3, 2], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
)];
let bytes = build_gguf("glm-dsa", &tensors);
let path = std::env::temp_dir().join(format!(
"ferrox_glm52_wk_b_test_{}.gguf",
std::process::id()
));
std::fs::write(&path, &bytes).unwrap();
let file = ferrox_gguf::GgufFile::open(&path).expect("synthetic GGUF must parse");
let heads = load_wk_b_transposed(&file, "blk.0.attn_k_b.weight", 1, 2, 3)
.expect("must load and transpose");
std::fs::remove_file(&path).ok();
assert_eq!(heads.len(), 1);
let applied_e0 = heads[0].apply(&[1.0, 0.0, 0.0]);
let applied_e1 = heads[0].apply(&[0.0, 1.0, 0.0]);
let applied_e2 = heads[0].apply(&[0.0, 0.0, 1.0]);
assert_eq!(applied_e0, vec![1.0, 2.0]);
assert_eq!(applied_e1, vec![3.0, 4.0]);
assert_eq!(applied_e2, vec![5.0, 6.0]);
}
}