Skip to main content

ferrox_models/
glm52_gguf_loader.rs

1//! Loads `ferrox-models::glm52_decoder` weights from a GLM-5.2 GGUF
2//! checkpoint, via `ferrox_gguf::TensorSource` -- the same trait
3//! `ferrox-models::loader`'s generic GQA path and
4//! `kimi_gguf_loader`'s Kimi K3 path both use. Follows
5//! `kimi_gguf_loader.rs`'s dedicated-loader pattern (a hand-written
6//! loader for an architecture whose MLA+DSA structure doesn't fit the
7//! generic GQA loader), not `loader.rs`'s generic path.
8//!
9//! Every tensor name/shape here is confirmed against real, inspectable
10//! upstream source, not guessed: `ggerganov/llama.cpp` PR #23346
11//! (DeepSeek-V3.2's DSA graph builder, `src/models/deepseek32.cpp`,
12//! which GLM-5.2's own model file is a fork of) and PR #25407
13//! (GLM-5.2's `indexer_types` diff on top,
14//! `src/models/glm-dsa.cpp::load_arch_tensors`'s real
15//! `create_tensor(tn(LLM_TENSOR_...))` calls) -- fetched live via
16//! `gh api -H "Accept: application/vnd.github.raw"
17//! repos/ggerganov/llama.cpp/contents/src/models/glm-dsa.cpp` (rather
18//! than `gh pr diff`, since PR #25407 doesn't touch most of these
19//! `create_tensor` calls -- they predate it, inherited unchanged from
20//! the DeepSeek-V3.2 fork point) and cross-checked against
21//! `src/llama-arch.cpp`'s real `LLM_TENSOR_NAMES` table for the exact
22//! on-disk name strings. This loader has NOT been run against a real
23//! GLM-5.2 GGUF file (~744B params, no feasible download/quant exists
24//! at a size this environment could hold) -- it is real, inspectable
25//! code built from real upstream evidence, tested here against a small
26//! synthetic on-disk fixture, the same rigor `kimi_gguf_loader.rs`
27//! documents for its own untested-against-a-real-file status.
28//!
29//! One real, non-obvious fact this loader has to account for:
30//! **the real on-disk `attn_k_b`/`attn_v_b` tensors are separate
31//! per-head 3D tensors, not a combined 2D `kv_b_proj` the way Kimi K3's
32//! real checkpoint stores it.** `attn_k_b`'s real ggml `ne[]` is
33//! `[qk_nope_head_dim, kv_lora_rank, n_head]` -- llama.cpp applies it in
34//! the "absorbed" direction (`ggml_mul_mat(wk_b, q_nope)`, projecting
35//! the *query* into compressed space, a compute optimization) rather
36//! than decompressing K directly. This loader instead TRANSPOSES each
37//! head's slice at load time into `[qk_nope_head_dim, kv_lora_rank]`
38//! (the direct decompression direction), so `glm_dsa`'s attention
39//! forward pass can reuse the same un-absorbed math
40//! `ferrox_models::mla` already has via `causal_mla_attention_sparse`,
41//! instead of implementing a second, absorbed-computation attention
42//! primitive purely for compute-efficiency parity with llama.cpp (which
43//! this CPU reference implementation doesn't need). `attn_v_b`'s real
44//! `ne[]` is `[kv_lora_rank, v_head_dim, n_head]`, which is *already*
45//! in the needed decompression direction per head (`[v_head_dim,
46//! kv_lora_rank]`) -- no transpose needed there. The transpose is only
47//! implemented for F32/BF16 (this loader dequantizes to f32 first, then
48//! transposes elementwise) -- a real, disclosed gap for quantized
49//! `attn_k_b` tensors specifically, not silently wrong: `load_wk_b_head`
50//! returns a clear `LoadError::UnsupportedDtype` for any other quant
51//! kind rather than guessing.
52
53use ferrox_core::tensor::Tensor;
54use ferrox_core::weight_matrix::quant_kind_for;
55use ferrox_core::weight_matrix::{WeightBytes, WeightMatrix};
56use ferrox_gguf::{GgmlType, TensorSource};
57
58use crate::glm_dsa::{Glm52AttnWeights, Glm52MlaConfig, IndexerConfig, IndexerWeights};
59use crate::loader::LoadError;
60use crate::loader::{find_info, load_f32_vec, load_weight_matrix, split_expert_tensor};
61
62/// Splits `blk.N.attn_k_b.weight` (real on-disk ggml `ne[]` =
63/// `[qk_nope_head_dim, kv_lora_rank, n_head]`, i.e. per head, physically
64/// `kv_lora_rank` rows of `qk_nope_head_dim` floats each) into per-head
65/// `WeightMatrix`es, TRANSPOSED into `[qk_nope_head_dim, kv_lora_rank]`
66/// -- see module doc comment for why. F32/BF16 only (a real, disclosed
67/// gap for quantized `attn_k_b`, not a silent wrong-shape read).
68fn load_wk_b_transposed(
69    file: &impl TensorSource,
70    name: &str,
71    n_head: usize,
72    qk_nope_head_dim: usize,
73    kv_lora_rank: usize,
74) -> Result<Vec<WeightMatrix>, LoadError> {
75    let info = find_info(file, name)?;
76    if info.shape.len() != 3
77        || info.shape[0] as usize != qk_nope_head_dim
78        || info.shape[1] as usize != kv_lora_rank
79        || info.shape[2] as usize != n_head
80    {
81        return Err(LoadError::UnsupportedDtype(
82            format!(
83                "{name} (expected ne=[{qk_nope_head_dim}, {kv_lora_rank}, {n_head}], got {:?})",
84                info.shape
85            ),
86            info.dtype,
87        ));
88    }
89    if !matches!(info.dtype, GgmlType::F32 | GgmlType::F16 | GgmlType::BF16) {
90        return Err(LoadError::UnsupportedDtype(name.to_string(), info.dtype));
91    }
92    let all = load_f32_vec(file, name)?;
93    let per_head = kv_lora_rank * qk_nope_head_dim;
94    Ok((0..n_head)
95        .map(|h| {
96            let head_raw = &all[h * per_head..(h + 1) * per_head]; // [kv_lora_rank, qk_nope_head_dim] row-major
97            let mut transposed = vec![0f32; per_head]; // [qk_nope_head_dim, kv_lora_rank] row-major
98            for row in 0..kv_lora_rank {
99                for col in 0..qk_nope_head_dim {
100                    transposed[col * kv_lora_rank + row] = head_raw[row * qk_nope_head_dim + col];
101                }
102            }
103            WeightMatrix::F32(Tensor::new(
104                transposed,
105                vec![qk_nope_head_dim, kv_lora_rank],
106            ))
107        })
108        .collect())
109}
110
111/// Splits `blk.N.attn_v_b.weight` (real on-disk ggml `ne[]` =
112/// `[kv_lora_rank, v_head_dim, n_head]`) into per-head `WeightMatrix`es
113/// -- already in the needed decompression direction (`[v_head_dim,
114/// kv_lora_rank]` per head), no transpose required, unlike `attn_k_b`.
115fn load_wv_b(
116    file: &impl TensorSource,
117    name: &str,
118    n_head: usize,
119    kv_lora_rank: usize,
120    v_head_dim: usize,
121) -> Result<Vec<WeightMatrix>, LoadError> {
122    let info = find_info(file, name)?;
123    if info.shape.len() != 3
124        || info.shape[0] as usize != kv_lora_rank
125        || info.shape[1] as usize != v_head_dim
126        || info.shape[2] as usize != n_head
127    {
128        return Err(LoadError::UnsupportedDtype(
129            format!(
130                "{name} (expected ne=[{kv_lora_rank}, {v_head_dim}, {n_head}], got {:?})",
131                info.shape
132            ),
133            info.dtype,
134        ));
135    }
136    match info.dtype {
137        GgmlType::F32 | GgmlType::F16 | GgmlType::BF16 => {
138            let all = load_f32_vec(file, name)?;
139            let per_head = kv_lora_rank * v_head_dim;
140            Ok((0..n_head)
141                .map(|h| {
142                    WeightMatrix::F32(Tensor::new(
143                        all[h * per_head..(h + 1) * per_head].to_vec(),
144                        vec![v_head_dim, kv_lora_rank],
145                    ))
146                })
147                .collect())
148        }
149        other => match quant_kind_for(other) {
150            Some(kind) => {
151                let (mmap, full_range) = file.tensor_mapped_range(name)?;
152                let bytes_per_head = (full_range.end - full_range.start) / n_head;
153                Ok((0..n_head)
154                    .map(|h| WeightMatrix::Quantized {
155                        data: WeightBytes::Mapped {
156                            mmap: std::sync::Arc::clone(&mmap),
157                            range: (full_range.start + h * bytes_per_head)
158                                ..(full_range.start + (h + 1) * bytes_per_head),
159                        },
160                        rows: v_head_dim,
161                        cols: kv_lora_rank,
162                        kind,
163                    })
164                    .collect())
165            }
166            None => Err(LoadError::UnsupportedDtype(name.to_string(), other)),
167        },
168    }
169}
170
171/// Real per-layer/global hyperparameters needed to load a GLM-5.2 GGUF
172/// file -- see docs/MODELS.md's "GLM-5.2 (Z.ai)" section for the
173/// real, confirmed values (78 layers, `hidden_size`=6144, etc.).
174pub struct Glm52GgufHparams {
175    pub hidden_dim: usize,
176    pub num_heads: usize,
177    pub q_lora_rank: usize,
178    pub kv_lora_rank: usize,
179    pub qk_nope_head_dim: usize,
180    pub qk_rope_head_dim: usize,
181    pub v_head_dim: usize,
182    pub rope_theta: f32,
183    pub indexer_n_heads: usize,
184    pub indexer_head_dim: usize,
185    pub indexer_rope_dim: usize,
186    pub indexer_top_k: usize,
187    pub dense_ffn_dim: usize,
188    pub moe_ffn_dim: usize,
189    pub n_experts: usize,
190    pub n_shared_experts: usize,
191}
192
193/// Loads one GLM-5.2 attention layer's weights. `is_full_indexer_layer`
194/// controls whether the real (`TENSOR_NOT_REQUIRED`-flagged, but always
195/// present for a "full" layer in a real checkpoint) indexer tensors are
196/// read at all -- "shared" layers carry no indexer weights of their own
197/// (see `glm_dsa`'s module doc comment point 1).
198pub fn load_glm52_attn(
199    file: &impl TensorSource,
200    hp: &Glm52GgufHparams,
201    layer_idx: usize,
202    is_full_indexer_layer: bool,
203) -> Result<Glm52AttnWeights, LoadError> {
204    let l = layer_idx;
205    let q_head_dim = hp.qk_nope_head_dim + hp.qk_rope_head_dim;
206
207    let q_a_proj = load_weight_matrix(file, &format!("blk.{l}.attn_q_a.weight"))?;
208    assert_eq!(
209        q_a_proj.rows(),
210        hp.q_lora_rank,
211        "blk.{l}.attn_q_a.weight row count"
212    );
213    assert_eq!(
214        q_a_proj.cols(),
215        hp.hidden_dim,
216        "blk.{l}.attn_q_a.weight col count"
217    );
218
219    let q_b_proj = load_weight_matrix(file, &format!("blk.{l}.attn_q_b.weight"))?;
220    assert_eq!(
221        q_b_proj.rows(),
222        hp.num_heads * q_head_dim,
223        "blk.{l}.attn_q_b.weight row count"
224    );
225
226    let kv_a_proj_with_mqa = load_weight_matrix(file, &format!("blk.{l}.attn_kv_a_mqa.weight"))?;
227    assert_eq!(
228        kv_a_proj_with_mqa.rows(),
229        hp.kv_lora_rank + hp.qk_rope_head_dim,
230        "blk.{l}.attn_kv_a_mqa.weight row count"
231    );
232
233    let wk_b = load_wk_b_transposed(
234        file,
235        &format!("blk.{l}.attn_k_b.weight"),
236        hp.num_heads,
237        hp.qk_nope_head_dim,
238        hp.kv_lora_rank,
239    )?;
240    let wv_b = load_wv_b(
241        file,
242        &format!("blk.{l}.attn_v_b.weight"),
243        hp.num_heads,
244        hp.kv_lora_rank,
245        hp.v_head_dim,
246    )?;
247
248    let o_proj = load_weight_matrix(file, &format!("blk.{l}.attn_output.weight"))?;
249    assert_eq!(
250        o_proj.rows(),
251        hp.hidden_dim,
252        "blk.{l}.attn_output.weight row count"
253    );
254    assert_eq!(
255        o_proj.cols(),
256        hp.num_heads * hp.v_head_dim,
257        "blk.{l}.attn_output.weight col count"
258    );
259
260    let indexer = if is_full_indexer_layer {
261        let k_norm_weight = load_f32_vec(file, &format!("blk.{l}.indexer.k_norm.weight"))?;
262        let k_norm_bias = load_f32_vec(file, &format!("blk.{l}.indexer.k_norm.bias"))?;
263        let proj = load_weight_matrix(file, &format!("blk.{l}.indexer.proj.weight"))?;
264        assert_eq!(
265            proj.rows(),
266            hp.indexer_n_heads,
267            "blk.{l}.indexer.proj.weight row count"
268        );
269        let attn_k = load_weight_matrix(file, &format!("blk.{l}.indexer.attn_k.weight"))?;
270        assert_eq!(
271            attn_k.rows(),
272            hp.indexer_head_dim,
273            "blk.{l}.indexer.attn_k.weight row count"
274        );
275        let attn_q_b = load_weight_matrix(file, &format!("blk.{l}.indexer.attn_q_b.weight"))?;
276        assert_eq!(
277            attn_q_b.rows(),
278            hp.indexer_n_heads * hp.indexer_head_dim,
279            "blk.{l}.indexer.attn_q_b.weight row count"
280        );
281        Some(IndexerWeights {
282            k_norm_weight,
283            k_norm_bias,
284            proj,
285            attn_k,
286            attn_q_b,
287        })
288    } else {
289        None
290    };
291
292    Ok(Glm52AttnWeights {
293        q_a_proj,
294        q_a_layernorm: load_f32_vec(file, &format!("blk.{l}.attn_q_a_norm.weight"))?,
295        q_b_proj,
296        kv_a_proj_with_mqa,
297        kv_a_layernorm: load_f32_vec(file, &format!("blk.{l}.attn_kv_a_norm.weight"))?,
298        wk_b,
299        wv_b,
300        o_proj,
301        indexer,
302    })
303}
304
305pub fn glm52_mla_config(hp: &Glm52GgufHparams) -> Glm52MlaConfig {
306    Glm52MlaConfig {
307        num_heads: hp.num_heads,
308        q_lora_rank: hp.q_lora_rank,
309        kv_lora_rank: hp.kv_lora_rank,
310        qk_nope_head_dim: hp.qk_nope_head_dim,
311        qk_rope_head_dim: hp.qk_rope_head_dim,
312        v_head_dim: hp.v_head_dim,
313        rope: crate::config::MlaRopeConfig {
314            theta: hp.rope_theta,
315        },
316    }
317}
318
319pub fn glm52_indexer_config(hp: &Glm52GgufHparams) -> IndexerConfig {
320    IndexerConfig {
321        n_heads: hp.indexer_n_heads,
322        head_dim: hp.indexer_head_dim,
323        rope_dim: hp.indexer_rope_dim,
324        top_k: hp.indexer_top_k,
325        rope_theta: hp.rope_theta,
326    }
327}
328
329/// Loads one dense leading layer's feed-forward block (real tensor
330/// names `blk.{bid}.ffn_{gate,down,up}` -- same convention every other
331/// architecture's dense FFN uses in this codebase).
332pub struct Glm52DenseFfnWeights {
333    pub gate_proj: WeightMatrix,
334    pub up_proj: WeightMatrix,
335    pub down_proj: WeightMatrix,
336}
337
338pub fn load_glm52_dense_ffn(
339    file: &impl TensorSource,
340    layer_idx: usize,
341) -> Result<Glm52DenseFfnWeights, LoadError> {
342    let l = layer_idx;
343    Ok(Glm52DenseFfnWeights {
344        gate_proj: load_weight_matrix(file, &format!("blk.{l}.ffn_gate.weight"))?,
345        up_proj: load_weight_matrix(file, &format!("blk.{l}.ffn_up.weight"))?,
346        down_proj: load_weight_matrix(file, &format!("blk.{l}.ffn_down.weight"))?,
347    })
348}
349
350/// One routed expert's gate/up/down weights (SwiGLU FFN) --
351/// `ferrox_moe::ExpertWeights`'s field names/order.
352pub struct Glm52MoeFfnWeights {
353    pub router_weight: WeightMatrix,
354    pub e_score_correction_bias: Vec<f32>,
355    pub experts: Vec<ferrox_moe::ExpertWeights>,
356    pub shared_expert: ferrox_moe::ExpertWeights,
357}
358
359pub fn load_glm52_moe_ffn(
360    file: &impl TensorSource,
361    hp: &Glm52GgufHparams,
362    layer_idx: usize,
363) -> Result<Glm52MoeFfnWeights, LoadError> {
364    let l = layer_idx;
365    let gate_exps =
366        split_expert_tensor(file, &format!("blk.{l}.ffn_gate_exps.weight"), hp.n_experts)?;
367    let down_exps =
368        split_expert_tensor(file, &format!("blk.{l}.ffn_down_exps.weight"), hp.n_experts)?;
369    let up_exps = split_expert_tensor(file, &format!("blk.{l}.ffn_up_exps.weight"), hp.n_experts)?;
370    let experts = gate_exps
371        .into_iter()
372        .zip(down_exps)
373        .zip(up_exps)
374        .map(|((gate, down), up)| ferrox_moe::ExpertWeights { gate, up, down })
375        .collect();
376
377    let shared_expert = ferrox_moe::ExpertWeights {
378        gate: load_weight_matrix(file, &format!("blk.{l}.ffn_gate_shexp.weight"))?,
379        up: load_weight_matrix(file, &format!("blk.{l}.ffn_up_shexp.weight"))?,
380        down: load_weight_matrix(file, &format!("blk.{l}.ffn_down_shexp.weight"))?,
381    };
382
383    Ok(Glm52MoeFfnWeights {
384        router_weight: load_weight_matrix(file, &format!("blk.{l}.ffn_gate_inp.weight"))?,
385        // `LLM_TENSOR_FFN_EXP_PROBS_B` is spelled `blk.%d.exp_probs_b`
386        // on disk -- no `ffn_` prefix (llama-arch.cpp:416,
387        // gguf-py/gguf/constants.py:1240). This asked for a name no real
388        // checkpoint carries, and the synthetic fixture below wrote the
389        // same wrong name, so the tests agreed with the bug.
390        e_score_correction_bias: load_f32_vec(file, &format!("blk.{l}.exp_probs_b.bias"))?,
391        experts,
392        shared_expert,
393    })
394}
395
396fn meta_u64(file: &impl TensorSource, key: &str) -> Result<u64, LoadError> {
397    file.metadata_u64(key)
398        .ok_or_else(|| LoadError::MissingHparam(key.to_string()))
399}
400
401fn meta_f32(file: &impl TensorSource, key: &str, default: f32) -> f32 {
402    file.metadata_f32(key).unwrap_or(default)
403}
404
405/// GGUF-side metadata needed alongside [`Glm52GgufHparams`] to build
406/// [`crate::engine::Glm52Engine`].
407pub struct Glm52FileMeta {
408    pub arch: String,
409    pub n_layer: usize,
410    pub leading_dense: usize,
411    pub rms_norm_eps: f32,
412    pub n_experts_active: usize,
413    pub moe_renormalize: bool,
414    pub routed_scaling_factor: f32,
415}
416
417/// Read GLM-5.2 / GLM4-family hparams from an opened GGUF (`glm-dsa`,
418/// `glm4`, `glm4moe`).
419pub fn read_glm52_hparams(
420    file: &impl TensorSource,
421) -> Result<(Glm52GgufHparams, Glm52FileMeta), LoadError> {
422    let arch = file
423        .metadata_str("general.architecture")
424        .ok_or_else(|| LoadError::MissingHparam("general.architecture".into()))?
425        .to_string();
426    if !matches!(arch.as_str(), "glm-dsa" | "glm4" | "glm4moe") {
427        return Err(LoadError::UnsupportedArchitecture(arch));
428    }
429    let p = |suffix: &str| format!("{arch}.{suffix}");
430    let n_layer = meta_u64(file, &p("block_count"))? as usize;
431    let hidden_dim = meta_u64(file, &p("embedding_length"))? as usize;
432    let dense_ffn_dim = meta_u64(file, &p("feed_forward_length"))? as usize;
433    let moe_ffn_dim = file
434        .metadata_u64(&p("expert_feed_forward_length"))
435        .unwrap_or(dense_ffn_dim as u64) as usize;
436    let n_heads = meta_u64(file, &p("attention.head_count"))? as usize;
437    let q_lora_rank = meta_u64(file, &p("attention.q_lora_rank"))? as usize;
438    let kv_lora_rank = meta_u64(file, &p("attention.kv_lora_rank"))? as usize;
439    let qk_nope_head_dim = meta_u64(file, &p("attention.qk_nope_head_dim"))? as usize;
440    let qk_rope_head_dim = meta_u64(file, &p("attention.qk_rope_head_dim"))? as usize;
441    let v_head_dim = file
442        .metadata_u64(&p("attention.v_head_dim"))
443        .or_else(|| file.metadata_u64(&p("attention.key_length")))
444        .unwrap_or(qk_nope_head_dim as u64) as usize;
445    let leading_dense = file
446        .metadata_u64(&p("leading_dense_block_count"))
447        .unwrap_or(0) as usize;
448    let n_experts = file.metadata_u64(&p("expert_count")).unwrap_or(0) as usize;
449    let n_shared_experts = file.metadata_u64(&p("expert_shared_count")).unwrap_or(1) as usize;
450    let n_experts_active = file.metadata_u64(&p("expert_used_count")).unwrap_or(8) as usize;
451    let indexer_n_heads = file
452        .metadata_u64(&p("attention.indexer_n_heads"))
453        .unwrap_or(4) as usize;
454    let indexer_head_dim = file
455        .metadata_u64(&p("attention.indexer_head_dim"))
456        .unwrap_or(128) as usize;
457    let indexer_top_k = file
458        .metadata_u64(&p("attention.indexer_top_k"))
459        .unwrap_or(2048) as usize;
460    let hp = Glm52GgufHparams {
461        hidden_dim,
462        num_heads: n_heads,
463        q_lora_rank,
464        kv_lora_rank,
465        qk_nope_head_dim,
466        qk_rope_head_dim,
467        v_head_dim,
468        rope_theta: meta_f32(file, &p("rope.freq_base"), 1_000_000.0),
469        indexer_n_heads,
470        indexer_head_dim,
471        indexer_rope_dim: qk_rope_head_dim,
472        indexer_top_k,
473        dense_ffn_dim,
474        moe_ffn_dim,
475        n_experts,
476        n_shared_experts,
477    };
478    let meta = Glm52FileMeta {
479        arch: arch.clone(),
480        n_layer,
481        leading_dense: leading_dense.min(n_layer),
482        rms_norm_eps: meta_f32(file, &p("attention.layer_norm_rms_epsilon"), 1e-5),
483        n_experts_active,
484        moe_renormalize: file
485            .metadata_u64(&p("expert_norm_topk_prob"))
486            .is_some_and(|v| v != 0),
487        routed_scaling_factor: meta_f32(file, &p("expert_routing_scale"), 2.5),
488    };
489    Ok((hp, meta))
490}
491
492fn is_full_indexer_layer(file: &impl TensorSource, layer_idx: usize) -> bool {
493    file.find_tensor(&format!("blk.{layer_idx}.indexer.proj.weight"))
494        .is_some()
495}
496
497fn is_dense_ffn_layer(file: &impl TensorSource, layer_idx: usize, leading_dense: usize) -> bool {
498    if layer_idx < leading_dense {
499        return true;
500    }
501    file.find_tensor(&format!("blk.{layer_idx}.ffn_gate.weight"))
502        .is_some()
503        && file
504            .find_tensor(&format!("blk.{layer_idx}.ffn_gate_inp.weight"))
505            .is_none()
506}
507
508fn load_embedding_tensor(
509    file: &impl TensorSource,
510    hidden_dim: usize,
511) -> Result<ferrox_core::tensor::Tensor, LoadError> {
512    let wm = load_weight_matrix(file, "token_embd.weight")?;
513    let vocab = wm.rows();
514    assert_eq!(wm.cols(), hidden_dim, "token_embd.weight col count");
515    let mut data = vec![0f32; vocab * hidden_dim];
516    for row in 0..vocab {
517        let r = wm.dequant_row(row);
518        data[row * hidden_dim..(row + 1) * hidden_dim].copy_from_slice(&r);
519    }
520    Ok(ferrox_core::tensor::Tensor::new(
521        data,
522        vec![vocab, hidden_dim],
523    ))
524}
525
526/// Load a GLM-5.2 / GLM4-family GGUF into [`crate::engine::Glm52Engine`].
527pub fn load_glm52_engine(
528    file: &impl TensorSource,
529) -> Result<crate::engine::Glm52Engine, LoadError> {
530    use crate::glm52_decoder::{
531        Glm52DecoderConfig, Glm52DecoderLayerWeights, Glm52DecoderWeights,
532        Glm52DenseFfnWeights as DecDenseFfn, Glm52LayerFfn, Glm52MoeFfnWeights as DecMoeFfn,
533    };
534
535    let (hp, meta) = read_glm52_hparams(file)?;
536    let embedding = load_embedding_tensor(file, hp.hidden_dim)?;
537    let final_norm_weight = load_f32_vec(file, "output_norm.weight")?;
538    let output_head = match load_weight_matrix(file, "output.weight") {
539        Ok(w) => w,
540        Err(_) => load_weight_matrix(file, "token_embd.weight")?,
541    };
542
543    let mut layers = Vec::with_capacity(meta.n_layer);
544    for layer_idx in 0..meta.n_layer {
545        let is_full = is_full_indexer_layer(file, layer_idx);
546        let attn = load_glm52_attn(file, &hp, layer_idx, is_full)?;
547        let ffn = if is_dense_ffn_layer(file, layer_idx, meta.leading_dense) {
548            let d = load_glm52_dense_ffn(file, layer_idx)?;
549            Glm52LayerFfn::Dense(Box::new(DecDenseFfn {
550                gate_proj: d.gate_proj,
551                up_proj: d.up_proj,
552                down_proj: d.down_proj,
553            }))
554        } else {
555            let m = load_glm52_moe_ffn(file, &hp, layer_idx)?;
556            Glm52LayerFfn::Moe(Box::new(DecMoeFfn {
557                router_weight: m.router_weight,
558                e_score_correction_bias: m.e_score_correction_bias,
559                experts: m.experts,
560                shared_expert: m.shared_expert,
561            }))
562        };
563        layers.push(Glm52DecoderLayerWeights {
564            attn_norm_weight: load_f32_vec(file, &format!("blk.{layer_idx}.attn_norm.weight"))?,
565            attn,
566            ffn_norm_weight: load_f32_vec(file, &format!("blk.{layer_idx}.ffn_norm.weight"))?,
567            ffn,
568            is_full_indexer_layer: is_full,
569        });
570    }
571
572    let weights = Glm52DecoderWeights {
573        embedding,
574        layers,
575        final_norm_weight,
576        output_head,
577    };
578    let cfg = Glm52DecoderConfig {
579        rms_norm_eps: meta.rms_norm_eps,
580        mla: glm52_mla_config(&hp),
581        indexer: glm52_indexer_config(&hp),
582        n_experts_active: meta.n_experts_active,
583        moe_renormalize: meta.moe_renormalize,
584        routed_scaling_factor: meta.routed_scaling_factor,
585    };
586    Ok(crate::engine::Glm52Engine { weights, cfg })
587}
588
589#[cfg(test)]
590mod tests {
591    use super::*;
592    use byteorder::{LittleEndian, WriteBytesExt};
593    use std::io::Write;
594
595    fn f32_bytes(values: &[f32]) -> Vec<u8> {
596        values.iter().flat_map(|v| v.to_le_bytes()).collect()
597    }
598
599    struct FixtureTensor {
600        name: String,
601        shape: Vec<u64>,
602        bytes: Vec<u8>,
603    }
604
605    fn f32_tensor(name: impl Into<String>, shape: Vec<u64>, values: Vec<f32>) -> FixtureTensor {
606        FixtureTensor {
607            name: name.into(),
608            shape,
609            bytes: f32_bytes(&values),
610        }
611    }
612
613    /// Builds a real, parseable on-disk GGUF file from a flat list of
614    /// tensors -- same builder pattern as
615    /// `kimi_gguf_loader::tests::build_gguf` (duplicated, not shared,
616    /// matching that module's own precedent), F32-only since none of
617    /// the hparams under test here are quantized-format-sensitive.
618    fn build_gguf(arch: &str, tensors: &[FixtureTensor]) -> Vec<u8> {
619        let mut buf = Vec::new();
620        buf.write_u32::<LittleEndian>(ferrox_gguf::GGUF_MAGIC)
621            .unwrap();
622        buf.write_u32::<LittleEndian>(3).unwrap(); // version
623        buf.write_u64::<LittleEndian>(tensors.len() as u64).unwrap();
624        buf.write_u64::<LittleEndian>(1).unwrap(); // kv_count: just general.architecture
625
626        let write_string = |buf: &mut Vec<u8>, s: &str| {
627            buf.write_u64::<LittleEndian>(s.len() as u64).unwrap();
628            buf.write_all(s.as_bytes()).unwrap();
629        };
630        write_string(&mut buf, "general.architecture");
631        buf.write_u32::<LittleEndian>(8).unwrap(); // type = string
632        write_string(&mut buf, arch);
633
634        let mut offset = 0u64;
635        let mut offsets = Vec::with_capacity(tensors.len());
636        for t in tensors {
637            write_string(&mut buf, &t.name);
638            buf.write_u32::<LittleEndian>(t.shape.len() as u32).unwrap();
639            for &d in t.shape.iter().rev() {
640                buf.write_u64::<LittleEndian>(d).unwrap();
641            }
642            buf.write_u32::<LittleEndian>(0).unwrap(); // dtype tag: F32
643            offsets.push(offset);
644            buf.write_u64::<LittleEndian>(offset).unwrap();
645            let padded = t.bytes.len().div_ceil(32) * 32;
646            offset += padded as u64;
647        }
648
649        while buf.len() % 32 != 0 {
650            buf.push(0);
651        }
652        let data_start = buf.len();
653        for (t, &off) in tensors.iter().zip(offsets.iter()) {
654            let want_len = data_start + off as usize;
655            while buf.len() < want_len {
656                buf.push(0);
657            }
658            buf.extend_from_slice(&t.bytes);
659            while buf.len() % 32 != 0 {
660                buf.push(0);
661            }
662        }
663        buf
664    }
665
666    struct Dims {
667        hidden_dim: usize,
668        num_heads: usize,
669        q_lora_rank: usize,
670        kv_lora_rank: usize,
671        qk_nope_head_dim: usize,
672        qk_rope_head_dim: usize,
673        v_head_dim: usize,
674        indexer_n_heads: usize,
675        indexer_head_dim: usize,
676        dense_ffn_dim: usize,
677        moe_ffn_dim: usize,
678        n_experts: usize,
679        n_shared_experts: usize,
680    }
681
682    fn push_layer_tensors(
683        tensors: &mut Vec<FixtureTensor>,
684        l: usize,
685        is_full: bool,
686        is_dense: bool,
687        d: &Dims,
688    ) {
689        let h = d.hidden_dim;
690        let q_head_dim = d.qk_nope_head_dim + d.qk_rope_head_dim;
691
692        tensors.push(f32_tensor(
693            format!("blk.{l}.attn_norm.weight"),
694            vec![h as u64],
695            vec![1.0; h],
696        ));
697        tensors.push(f32_tensor(
698            format!("blk.{l}.ffn_norm.weight"),
699            vec![h as u64],
700            vec![1.0; h],
701        ));
702        tensors.push(f32_tensor(
703            format!("blk.{l}.attn_q_a_norm.weight"),
704            vec![d.q_lora_rank as u64],
705            vec![1.0; d.q_lora_rank],
706        ));
707        tensors.push(f32_tensor(
708            format!("blk.{l}.attn_kv_a_norm.weight"),
709            vec![d.kv_lora_rank as u64],
710            vec![1.0; d.kv_lora_rank],
711        ));
712        tensors.push(f32_tensor(
713            format!("blk.{l}.attn_q_a.weight"),
714            vec![d.q_lora_rank as u64, h as u64],
715            vec![0.02; d.q_lora_rank * h],
716        ));
717        tensors.push(f32_tensor(
718            format!("blk.{l}.attn_q_b.weight"),
719            vec![(d.num_heads * q_head_dim) as u64, d.q_lora_rank as u64],
720            vec![0.02; d.num_heads * q_head_dim * d.q_lora_rank],
721        ));
722        tensors.push(f32_tensor(
723            format!("blk.{l}.attn_kv_a_mqa.weight"),
724            vec![(d.kv_lora_rank + d.qk_rope_head_dim) as u64, h as u64],
725            vec![0.02; (d.kv_lora_rank + d.qk_rope_head_dim) * h],
726        ));
727        // Real ne = [qk_nope_head_dim, kv_lora_rank, n_head]; this
728        // builder reverses its `shape` arg once to produce the written
729        // raw ne, so the argument here must be the REVERSE of that.
730        tensors.push(f32_tensor(
731            format!("blk.{l}.attn_k_b.weight"),
732            vec![
733                d.num_heads as u64,
734                d.kv_lora_rank as u64,
735                d.qk_nope_head_dim as u64,
736            ],
737            vec![0.02; d.num_heads * d.kv_lora_rank * d.qk_nope_head_dim],
738        ));
739        tensors.push(f32_tensor(
740            format!("blk.{l}.attn_v_b.weight"),
741            vec![
742                d.num_heads as u64,
743                d.v_head_dim as u64,
744                d.kv_lora_rank as u64,
745            ],
746            vec![0.02; d.num_heads * d.v_head_dim * d.kv_lora_rank],
747        ));
748        tensors.push(f32_tensor(
749            format!("blk.{l}.attn_output.weight"),
750            vec![h as u64, (d.num_heads * d.v_head_dim) as u64],
751            vec![0.02; h * d.num_heads * d.v_head_dim],
752        ));
753
754        if is_full {
755            tensors.push(f32_tensor(
756                format!("blk.{l}.indexer.k_norm.weight"),
757                vec![d.indexer_head_dim as u64],
758                vec![1.0; d.indexer_head_dim],
759            ));
760            tensors.push(f32_tensor(
761                format!("blk.{l}.indexer.k_norm.bias"),
762                vec![d.indexer_head_dim as u64],
763                vec![0.0; d.indexer_head_dim],
764            ));
765            tensors.push(f32_tensor(
766                format!("blk.{l}.indexer.proj.weight"),
767                vec![d.indexer_n_heads as u64, h as u64],
768                vec![0.02; d.indexer_n_heads * h],
769            ));
770            tensors.push(f32_tensor(
771                format!("blk.{l}.indexer.attn_k.weight"),
772                vec![d.indexer_head_dim as u64, h as u64],
773                vec![0.02; d.indexer_head_dim * h],
774            ));
775            tensors.push(f32_tensor(
776                format!("blk.{l}.indexer.attn_q_b.weight"),
777                vec![
778                    (d.indexer_n_heads * d.indexer_head_dim) as u64,
779                    d.q_lora_rank as u64,
780                ],
781                vec![0.02; d.indexer_n_heads * d.indexer_head_dim * d.q_lora_rank],
782            ));
783        }
784
785        if is_dense {
786            for name in ["ffn_gate", "ffn_up"] {
787                tensors.push(f32_tensor(
788                    format!("blk.{l}.{name}.weight"),
789                    vec![d.dense_ffn_dim as u64, h as u64],
790                    vec![0.02; d.dense_ffn_dim * h],
791                ));
792            }
793            tensors.push(f32_tensor(
794                format!("blk.{l}.ffn_down.weight"),
795                vec![h as u64, d.dense_ffn_dim as u64],
796                vec![0.02; h * d.dense_ffn_dim],
797            ));
798        } else {
799            let ff = d.moe_ffn_dim;
800            let n = d.n_experts;
801            tensors.push(f32_tensor(
802                format!("blk.{l}.ffn_gate_inp.weight"),
803                vec![n as u64, h as u64],
804                vec![0.02; n * h],
805            ));
806            tensors.push(f32_tensor(
807                format!("blk.{l}.exp_probs_b.bias"),
808                vec![n as u64],
809                vec![0.0; n],
810            ));
811            tensors.push(f32_tensor(
812                format!("blk.{l}.ffn_gate_exps.weight"),
813                vec![n as u64, ff as u64, h as u64],
814                vec![0.02; h * ff * n],
815            ));
816            tensors.push(f32_tensor(
817                format!("blk.{l}.ffn_down_exps.weight"),
818                vec![n as u64, h as u64, ff as u64],
819                vec![0.02; ff * h * n],
820            ));
821            tensors.push(f32_tensor(
822                format!("blk.{l}.ffn_up_exps.weight"),
823                vec![n as u64, ff as u64, h as u64],
824                vec![0.02; h * ff * n],
825            ));
826            let shexp_dim = ff * d.n_shared_experts;
827            tensors.push(f32_tensor(
828                format!("blk.{l}.ffn_gate_shexp.weight"),
829                vec![shexp_dim as u64, h as u64],
830                vec![0.02; shexp_dim * h],
831            ));
832            tensors.push(f32_tensor(
833                format!("blk.{l}.ffn_down_shexp.weight"),
834                vec![h as u64, shexp_dim as u64],
835                vec![0.02; h * shexp_dim],
836            ));
837            tensors.push(f32_tensor(
838                format!("blk.{l}.ffn_up_shexp.weight"),
839                vec![shexp_dim as u64, h as u64],
840                vec![0.02; shexp_dim * h],
841            ));
842        }
843    }
844
845    fn dims() -> Dims {
846        Dims {
847            hidden_dim: 8,
848            num_heads: 2,
849            q_lora_rank: 6,
850            kv_lora_rank: 4,
851            qk_nope_head_dim: 4,
852            qk_rope_head_dim: 4,
853            v_head_dim: 3,
854            indexer_n_heads: 2,
855            indexer_head_dim: 4,
856            dense_ffn_dim: 5,
857            moe_ffn_dim: 4,
858            n_experts: 3,
859            n_shared_experts: 1,
860        }
861    }
862
863    fn hp_from(d: &Dims) -> Glm52GgufHparams {
864        Glm52GgufHparams {
865            hidden_dim: d.hidden_dim,
866            num_heads: d.num_heads,
867            q_lora_rank: d.q_lora_rank,
868            kv_lora_rank: d.kv_lora_rank,
869            qk_nope_head_dim: d.qk_nope_head_dim,
870            qk_rope_head_dim: d.qk_rope_head_dim,
871            v_head_dim: d.v_head_dim,
872            rope_theta: 8_000_000.0,
873            indexer_n_heads: d.indexer_n_heads,
874            indexer_head_dim: d.indexer_head_dim,
875            indexer_rope_dim: d.qk_rope_head_dim,
876            indexer_top_k: 2,
877            dense_ffn_dim: d.dense_ffn_dim,
878            moe_ffn_dim: d.moe_ffn_dim,
879            n_experts: d.n_experts,
880            n_shared_experts: d.n_shared_experts,
881        }
882    }
883
884    #[test]
885    fn loads_a_full_indexer_dense_layer_and_a_shared_indexer_moe_layer() {
886        let d = dims();
887        let mut tensors: Vec<FixtureTensor> = Vec::new();
888        push_layer_tensors(&mut tensors, 0, true, true, &d);
889        push_layer_tensors(&mut tensors, 1, false, false, &d);
890
891        let bytes = build_gguf("glm-dsa", &tensors);
892        let path = std::env::temp_dir().join(format!(
893            "ferrox_glm52_gguf_test_{}.gguf",
894            std::process::id()
895        ));
896        std::fs::write(&path, &bytes).unwrap();
897        let file = ferrox_gguf::GgufFile::open(&path).expect("synthetic GGUF must parse");
898
899        let hp = hp_from(&d);
900
901        let layer0 = load_glm52_attn(&file, &hp, 0, true).expect("full-indexer layer must load");
902        assert!(layer0.indexer.is_some());
903        assert_eq!(layer0.wk_b.len(), d.num_heads);
904        assert_eq!(layer0.wk_b[0].rows(), d.qk_nope_head_dim);
905        assert_eq!(layer0.wk_b[0].cols(), d.kv_lora_rank);
906        assert_eq!(layer0.wv_b[0].rows(), d.v_head_dim);
907        assert_eq!(layer0.wv_b[0].cols(), d.kv_lora_rank);
908        let dense0 = load_glm52_dense_ffn(&file, 0).expect("dense FFN must load");
909        assert_eq!(dense0.gate_proj.rows(), d.dense_ffn_dim);
910
911        let layer1 = load_glm52_attn(&file, &hp, 1, false).expect("shared-indexer layer must load");
912        assert!(
913            layer1.indexer.is_none(),
914            "a \"shared\" layer must not load its own indexer weights"
915        );
916        let moe1 = load_glm52_moe_ffn(&file, &hp, 1).expect("MoE FFN must load");
917        assert_eq!(moe1.experts.len(), d.n_experts);
918        assert_eq!(moe1.e_score_correction_bias.len(), d.n_experts);
919
920        std::fs::remove_file(&path).ok();
921    }
922
923    #[test]
924    fn wk_b_transpose_matches_hand_computed_values() {
925        // n_head=1, qk_nope_head_dim=2, kv_lora_rank=3. Real on-disk
926        // layout (ne=[2,3,1], i.e. 3 rows of 2 floats): row0=[1,2],
927        // row1=[3,4], row2=[5,6]. Transposed [2,3] result must be
928        // row0=[1,3,5], row1=[2,4,6].
929        let tensors = vec![f32_tensor(
930            "blk.0.attn_k_b.weight",
931            vec![1, 3, 2], // reversed real ne=[2,3,1]
932            vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
933        )];
934        let bytes = build_gguf("glm-dsa", &tensors);
935        let path = std::env::temp_dir().join(format!(
936            "ferrox_glm52_wk_b_test_{}.gguf",
937            std::process::id()
938        ));
939        std::fs::write(&path, &bytes).unwrap();
940        let file = ferrox_gguf::GgufFile::open(&path).expect("synthetic GGUF must parse");
941
942        let heads = load_wk_b_transposed(&file, "blk.0.attn_k_b.weight", 1, 2, 3)
943            .expect("must load and transpose");
944        std::fs::remove_file(&path).ok();
945
946        assert_eq!(heads.len(), 1);
947        let applied_e0 = heads[0].apply(&[1.0, 0.0, 0.0]);
948        let applied_e1 = heads[0].apply(&[0.0, 1.0, 0.0]);
949        let applied_e2 = heads[0].apply(&[0.0, 0.0, 1.0]);
950        // Column c of the transposed [2,3] matrix (applying basis
951        // vector e_c) must equal row c of the original [3,2] input.
952        assert_eq!(applied_e0, vec![1.0, 2.0]);
953        assert_eq!(applied_e1, vec![3.0, 4.0]);
954        assert_eq!(applied_e2, vec![5.0, 6.0]);
955    }
956}