Skip to main content

ferrox_models/
mla_gguf_loader.rs

1//! DeepSeek-2 / Mistral-4 GGUF → [`crate::engine::MlaEngine`].
2//!
3//! Tensor names follow llama.cpp `deepseek2` / `mistral4` (same graph):
4//! `blk.{i}.attn_q_a|attn_q_b|attn_kv_a_mqa|attn_kv_b|attn_output` plus
5//! optional `attn_q_a_norm` / `attn_kv_a_norm`. Dense FFN:
6//! `ffn_{gate,up,down}`. MoE after `leading_dense_block_count` uses
7//! `ffn_gate_inp` + packed `ffn_{gate,up,down}_exps` + shared
8//! `ffn_{gate,up,down}_shexp` (fail-closed if any are missing).
9//!
10//! `use_output_gate` is off (classic DeepSeek-2). RoPE uses interleaved
11//! Norm layout via [`crate::config::MlaRopeConfig`].
12
13use ferrox_gguf::TensorSource;
14use ferrox_moe::GatingFunction;
15
16use crate::config::{MlaConfig, MlaRopeConfig};
17use crate::engine::{
18    MlaDenseFfn, MlaEngine, MlaLayerFfn, MlaLayerWeights, MlaMoeFfn, MlaMoeRuntime,
19};
20use crate::loader::LoadError;
21use crate::loader::{load_f32_vec, load_weight_matrix, split_expert_tensor};
22use crate::mla::MlaAttnWeights;
23
24/// Hyperparameters read from `{arch}.*` GGUF metadata.
25#[derive(Debug, Clone)]
26pub struct Deepseek2Hparams {
27    pub arch: String,
28    pub n_layer: usize,
29    pub hidden_dim: usize,
30    pub ffn_dim: usize,
31    pub n_heads: usize,
32    pub q_lora_rank: usize,
33    pub kv_lora_rank: usize,
34    pub qk_nope_head_dim: usize,
35    pub qk_rope_head_dim: usize,
36    pub v_head_dim: usize,
37    pub rms_norm_eps: f32,
38    pub rope_theta: f32,
39    /// Layers `[0, leading_dense)` use dense SwiGLU; rest require MoE.
40    pub leading_dense_block_count: usize,
41    pub n_expert: usize,
42    pub n_expert_used: usize,
43    pub n_shared_experts: usize,
44    pub expert_ffn_dim: usize,
45    pub gating: GatingFunction,
46    pub norm_topk_prob: bool,
47    pub expert_weights_scale: f32,
48}
49
50fn meta_u64(file: &impl TensorSource, key: &str) -> Result<u64, LoadError> {
51    file.metadata_u64(key)
52        .ok_or_else(|| LoadError::MissingHparam(key.to_string()))
53}
54
55fn meta_f32(file: &impl TensorSource, key: &str, default: f32) -> f32 {
56    file.metadata_f32(key).unwrap_or(default)
57}
58
59/// Read DeepSeek-2 / Mistral-4 hparams from an opened GGUF.
60pub fn read_deepseek2_hparams(file: &impl TensorSource) -> Result<Deepseek2Hparams, LoadError> {
61    let arch = file
62        .metadata_str("general.architecture")
63        .ok_or_else(|| LoadError::MissingHparam("general.architecture".into()))?
64        .to_string();
65    if arch != "deepseek2" && arch != "mistral4" {
66        return Err(LoadError::UnsupportedArchitecture(arch));
67    }
68    let p = |suffix: &str| format!("{arch}.{suffix}");
69    let n_layer = meta_u64(file, &p("block_count"))? as usize;
70    let hidden_dim = meta_u64(file, &p("embedding_length"))? as usize;
71    let ffn_dim = meta_u64(file, &p("feed_forward_length"))? as usize;
72    let n_heads = meta_u64(file, &p("attention.head_count"))? as usize;
73    let q_lora_rank = meta_u64(file, &p("attention.q_lora_rank"))? as usize;
74    let kv_lora_rank = meta_u64(file, &p("attention.kv_lora_rank"))? as usize;
75    let qk_nope_head_dim = meta_u64(file, &p("attention.qk_nope_head_dim"))? as usize;
76    let qk_rope_head_dim = meta_u64(file, &p("attention.qk_rope_head_dim"))? as usize;
77    let v_head_dim = meta_u64(file, &p("attention.v_head_dim"))
78        .or_else(|_| meta_u64(file, &p("attention.key_length")))
79        .unwrap_or(qk_nope_head_dim as u64) as usize;
80    let leading_dense = file
81        .metadata_u64(&p("leading_dense_block_count"))
82        .unwrap_or(n_layer as u64) as usize;
83    let n_expert = file.metadata_u64(&p("expert_count")).unwrap_or(0) as usize;
84    let n_expert_used = file
85        .metadata_u64(&p("expert_used_count"))
86        .unwrap_or(if n_expert > 0 { 6 } else { 0 }) as usize;
87    let n_shared_experts = file.metadata_u64(&p("expert_shared_count")).unwrap_or(1) as usize;
88    let expert_ffn_dim = file
89        .metadata_u64(&p("expert_feed_forward_length"))
90        .unwrap_or(ffn_dim as u64) as usize;
91    let rms_norm_eps = meta_f32(file, &p("attention.layer_norm_rms_epsilon"), 1e-6);
92    let rope_theta = meta_f32(file, &p("rope.freq_base"), 10000.0);
93    // llama.cpp deepseek2: default Softmax unless expert_gating_func set
94    // (1=softmax, 2=sigmoid); special-case GLM 4.7 Lite sigmoid when absent.
95    let gating = match file.metadata_u64(&p("expert_gating_func")) {
96        Some(2) => GatingFunction::Sigmoid,
97        Some(1) => GatingFunction::Softmax,
98        _ if (n_layer == 47 || n_layer == 48)
99            && file
100                .find_tensor("token_embd.weight")
101                .map(|t| t.shape.last().copied().unwrap_or(0) == 154880)
102                .unwrap_or(false) =>
103        {
104            GatingFunction::Sigmoid
105        }
106        _ => GatingFunction::Softmax,
107    };
108    let norm_topk_prob = file
109        .metadata_u64(&p("expert_weights_norm"))
110        .map(|v| v != 0)
111        .unwrap_or(true);
112    let expert_weights_scale = meta_f32(file, &p("expert_weights_scale"), 1.0);
113    Ok(Deepseek2Hparams {
114        arch,
115        n_layer,
116        hidden_dim,
117        ffn_dim,
118        n_heads,
119        q_lora_rank,
120        kv_lora_rank,
121        qk_nope_head_dim,
122        qk_rope_head_dim,
123        v_head_dim,
124        rms_norm_eps,
125        rope_theta,
126        leading_dense_block_count: leading_dense.min(n_layer),
127        n_expert,
128        n_expert_used: n_expert_used.min(n_expert.max(1)),
129        n_shared_experts: n_shared_experts.max(1),
130        expert_ffn_dim,
131        gating,
132        norm_topk_prob,
133        expert_weights_scale,
134    })
135}
136
137fn load_f32_vec_optional(
138    file: &impl TensorSource,
139    name: &str,
140) -> Result<Option<Vec<f32>>, LoadError> {
141    if file.find_tensor(name).is_none() {
142        return Ok(None);
143    }
144    Ok(Some(load_f32_vec(file, name)?))
145}
146
147fn load_mla_attn(
148    file: &impl TensorSource,
149    layer_idx: usize,
150    hp: &Deepseek2Hparams,
151) -> Result<MlaAttnWeights, LoadError> {
152    let l = layer_idx;
153    let q_head_dim = hp.qk_nope_head_dim + hp.qk_rope_head_dim;
154    let q_a_proj = load_weight_matrix(file, &format!("blk.{l}.attn_q_a.weight"))?;
155    let q_b_proj = load_weight_matrix(file, &format!("blk.{l}.attn_q_b.weight"))?;
156    let kv_a = load_weight_matrix(file, &format!("blk.{l}.attn_kv_a_mqa.weight"))?;
157    let o_proj = load_weight_matrix(file, &format!("blk.{l}.attn_output.weight"))?;
158
159    // Prefer combined `attn_kv_b`; else refuse split k_b/v_b until concat lands.
160    let kv_b_proj = if file
161        .find_tensor(&format!("blk.{l}.attn_kv_b.weight"))
162        .is_some()
163    {
164        load_weight_matrix(file, &format!("blk.{l}.attn_kv_b.weight"))?
165    } else {
166        return Err(LoadError::Gguf(ferrox_gguf::GgufError::TensorNotFound(
167            format!(
168                "blk.{l}.attn_kv_b.weight (split attn_k_b/attn_v_b not wired for MlaEngine yet)"
169            ),
170        )));
171    };
172
173    let q_a_ln = load_f32_vec_optional(file, &format!("blk.{l}.attn_q_a_norm.weight"))?
174        .unwrap_or_else(|| vec![1.0; hp.q_lora_rank]);
175    let kv_a_ln = load_f32_vec_optional(file, &format!("blk.{l}.attn_kv_a_norm.weight"))?
176        .unwrap_or_else(|| vec![1.0; hp.kv_lora_rank]);
177
178    let _ = (q_head_dim,);
179    Ok(MlaAttnWeights {
180        q_a_proj,
181        q_a_layernorm: q_a_ln,
182        q_b_proj,
183        kv_a_proj_with_mqa: kv_a,
184        kv_a_layernorm: kv_a_ln,
185        kv_b_proj,
186        o_proj,
187        g_proj: None,
188    })
189}
190
191fn require_tensor(file: &impl TensorSource, name: &str) -> Result<(), LoadError> {
192    if file.find_tensor(name).is_none() {
193        return Err(LoadError::Gguf(ferrox_gguf::GgufError::TensorNotFound(
194            name.to_string(),
195        )));
196    }
197    Ok(())
198}
199
200fn load_dense_ffn(file: &impl TensorSource, layer_idx: usize) -> Result<MlaDenseFfn, LoadError> {
201    let l = layer_idx;
202    Ok(MlaDenseFfn {
203        gate: load_weight_matrix(file, &format!("blk.{l}.ffn_gate.weight"))?,
204        up: load_weight_matrix(file, &format!("blk.{l}.ffn_up.weight"))?,
205        down: load_weight_matrix(file, &format!("blk.{l}.ffn_down.weight"))?,
206    })
207}
208
209fn load_moe_ffn(
210    file: &impl TensorSource,
211    layer_idx: usize,
212    hp: &Deepseek2Hparams,
213) -> Result<MlaMoeFfn, LoadError> {
214    let l = layer_idx;
215    // Fail-closed: every MoE tensor must be present (no silent dense fallback).
216    for name in [
217        format!("blk.{l}.ffn_gate_inp.weight"),
218        format!("blk.{l}.ffn_gate_exps.weight"),
219        format!("blk.{l}.ffn_up_exps.weight"),
220        format!("blk.{l}.ffn_down_exps.weight"),
221        format!("blk.{l}.ffn_gate_shexp.weight"),
222        format!("blk.{l}.ffn_up_shexp.weight"),
223        format!("blk.{l}.ffn_down_shexp.weight"),
224    ] {
225        require_tensor(file, &name)?;
226    }
227    let gate_exps =
228        split_expert_tensor(file, &format!("blk.{l}.ffn_gate_exps.weight"), hp.n_expert)?;
229    let up_exps = split_expert_tensor(file, &format!("blk.{l}.ffn_up_exps.weight"), hp.n_expert)?;
230    let down_exps =
231        split_expert_tensor(file, &format!("blk.{l}.ffn_down_exps.weight"), hp.n_expert)?;
232    let experts = gate_exps
233        .into_iter()
234        .zip(up_exps)
235        .zip(down_exps)
236        .map(|((gate, up), down)| ferrox_moe::ExpertWeights { gate, up, down })
237        .collect();
238    let shared_expert = ferrox_moe::ExpertWeights {
239        gate: load_weight_matrix(file, &format!("blk.{l}.ffn_gate_shexp.weight"))?,
240        up: load_weight_matrix(file, &format!("blk.{l}.ffn_up_shexp.weight"))?,
241        down: load_weight_matrix(file, &format!("blk.{l}.ffn_down_shexp.weight"))?,
242    };
243    // See the note in `glm52_gguf_loader`: the on-disk name has no
244    // `ffn_` prefix. Optional here on purpose -- llama.cpp declares it
245    // TENSOR_NOT_REQUIRED for `deepseek2`, which also covers V2-era
246    // checkpoints with no routing bias at all -- which is exactly why the
247    // wrong name was silent rather than a load error, and a real
248    // DeepSeek-V3 checkpoint routed with its bias dropped.
249    let exp_probs_bias = load_f32_vec_optional(file, &format!("blk.{l}.exp_probs_b.bias"))?;
250    Ok(MlaMoeFfn {
251        router: load_weight_matrix(file, &format!("blk.{l}.ffn_gate_inp.weight"))?,
252        experts,
253        shared_expert,
254        exp_probs_bias,
255    })
256}
257
258fn load_layer(
259    file: &impl TensorSource,
260    layer_idx: usize,
261    hp: &Deepseek2Hparams,
262) -> Result<MlaLayerWeights, LoadError> {
263    let l = layer_idx;
264    let ffn = if layer_idx < hp.leading_dense_block_count || hp.n_expert == 0 {
265        MlaLayerFfn::Dense(load_dense_ffn(file, layer_idx)?)
266    } else {
267        MlaLayerFfn::Moe(load_moe_ffn(file, layer_idx, hp)?)
268    };
269    Ok(MlaLayerWeights {
270        attn_norm: load_f32_vec(file, &format!("blk.{l}.attn_norm.weight"))?,
271        attn: load_mla_attn(file, layer_idx, hp)?,
272        ffn_norm: load_f32_vec(file, &format!("blk.{l}.ffn_norm.weight"))?,
273        ffn,
274    })
275}
276
277/// Load a DeepSeek-2 / Mistral-4 GGUF into [`MlaEngine`] (dense lead + MoE tail).
278pub fn load_mla_engine(file: &impl TensorSource) -> Result<MlaEngine, LoadError> {
279    let hp = read_deepseek2_hparams(file)?;
280    if hp.n_expert > 0 && hp.leading_dense_block_count >= hp.n_layer {
281        // Experts declared but every layer is still dense — ignore MoE.
282    } else if hp.n_expert > 0 && hp.n_expert_used == 0 {
283        return Err(LoadError::UnsupportedArchitecture(format!(
284            "{}: expert_count={} but expert_used_count is 0",
285            hp.arch, hp.n_expert
286        )));
287    }
288    if hp.n_layer == 0 {
289        return Err(LoadError::UnsupportedArchitecture(format!(
290            "{}: no layers to load",
291            hp.arch
292        )));
293    }
294
295    let embedding = if file.find_tensor("token_embd.weight").is_some() {
296        load_weight_matrix(file, "token_embd.weight")?
297    } else {
298        return Err(LoadError::Gguf(ferrox_gguf::GgufError::TensorNotFound(
299            "token_embd.weight".into(),
300        )));
301    };
302    let final_norm = load_f32_vec(file, "output_norm.weight")?;
303    let output_head = match load_weight_matrix(file, "output.weight") {
304        Ok(w) => w,
305        Err(_) => load_weight_matrix(file, "token_embd.weight")?,
306    };
307
308    let mut layers = Vec::with_capacity(hp.n_layer);
309    for i in 0..hp.n_layer {
310        layers.push(load_layer(file, i, &hp)?);
311    }
312    let has_moe = layers.iter().any(|l| matches!(l.ffn, MlaLayerFfn::Moe(_)));
313    let moe = if has_moe {
314        Some(MlaMoeRuntime {
315            n_experts_active: hp.n_expert_used,
316            gating: hp.gating,
317            norm_topk_prob: hp.norm_topk_prob,
318            expert_weights_scale: hp.expert_weights_scale,
319        })
320    } else {
321        None
322    };
323
324    Ok(MlaEngine {
325        embedding,
326        layers,
327        final_norm,
328        output_head,
329        mla_cfg: MlaConfig {
330            num_heads: hp.n_heads,
331            q_lora_rank: hp.q_lora_rank,
332            kv_lora_rank: hp.kv_lora_rank,
333            qk_nope_head_dim: hp.qk_nope_head_dim,
334            qk_rope_head_dim: hp.qk_rope_head_dim,
335            v_head_dim: hp.v_head_dim,
336            use_output_gate: false,
337            rope: Some(MlaRopeConfig {
338                theta: hp.rope_theta,
339            }),
340        },
341        rms_norm_eps: hp.rms_norm_eps,
342        hidden_dim: hp.hidden_dim,
343        moe,
344    })
345}
346
347#[cfg(test)]
348mod tests {
349    use super::*;
350    use crate::engine::Engine;
351    use byteorder::{LittleEndian, WriteBytesExt};
352    use ferrox_gguf::GgufFile;
353    use std::io::Write;
354
355    struct FixtureTensor {
356        name: String,
357        shape: Vec<u64>,
358        bytes: Vec<u8>,
359    }
360
361    fn f32_bytes(v: &[f32]) -> Vec<u8> {
362        let mut b = Vec::with_capacity(v.len() * 4);
363        for x in v {
364            b.write_f32::<LittleEndian>(*x).unwrap();
365        }
366        b
367    }
368
369    fn f32_tensor(name: &str, shape: Vec<u64>, values: Vec<f32>) -> FixtureTensor {
370        FixtureTensor {
371            name: name.into(),
372            shape,
373            bytes: f32_bytes(&values),
374        }
375    }
376
377    fn build_gguf(
378        arch: &str,
379        kv: &[(&str, u64)],
380        fkv: &[(&str, f32)],
381        tensors: &[FixtureTensor],
382    ) -> Vec<u8> {
383        let mut buf = Vec::new();
384        buf.write_u32::<LittleEndian>(ferrox_gguf::GGUF_MAGIC)
385            .unwrap();
386        buf.write_u32::<LittleEndian>(3).unwrap();
387        buf.write_u64::<LittleEndian>(tensors.len() as u64).unwrap();
388        // general.architecture + uint + float kvs
389        let kv_count = 1 + kv.len() + fkv.len();
390        buf.write_u64::<LittleEndian>(kv_count as u64).unwrap();
391
392        let write_string = |buf: &mut Vec<u8>, s: &str| {
393            buf.write_u64::<LittleEndian>(s.len() as u64).unwrap();
394            buf.write_all(s.as_bytes()).unwrap();
395        };
396        write_string(&mut buf, "general.architecture");
397        buf.write_u32::<LittleEndian>(8).unwrap();
398        write_string(&mut buf, arch);
399        for &(k, v) in kv {
400            write_string(&mut buf, k);
401            buf.write_u32::<LittleEndian>(10).unwrap(); // UINT64
402            buf.write_u64::<LittleEndian>(v).unwrap();
403        }
404        for &(k, v) in fkv {
405            write_string(&mut buf, k);
406            buf.write_u32::<LittleEndian>(6).unwrap(); // FLOAT32
407            buf.write_f32::<LittleEndian>(v).unwrap();
408        }
409
410        let mut offset = 0u64;
411        let mut offsets = Vec::with_capacity(tensors.len());
412        for t in tensors {
413            write_string(&mut buf, &t.name);
414            buf.write_u32::<LittleEndian>(t.shape.len() as u32).unwrap();
415            for &d in t.shape.iter().rev() {
416                buf.write_u64::<LittleEndian>(d).unwrap();
417            }
418            buf.write_u32::<LittleEndian>(0).unwrap();
419            offsets.push(offset);
420            buf.write_u64::<LittleEndian>(offset).unwrap();
421            offset += (t.bytes.len().div_ceil(32) * 32) as u64;
422        }
423        while buf.len() % 32 != 0 {
424            buf.push(0);
425        }
426        let data_start = buf.len();
427        for (t, &off) in tensors.iter().zip(offsets.iter()) {
428            while buf.len() < data_start + off as usize {
429                buf.push(0);
430            }
431            buf.extend_from_slice(&t.bytes);
432            while buf.len() % 32 != 0 {
433                buf.push(0);
434            }
435        }
436        buf
437    }
438
439    #[test]
440    fn load_synthetic_deepseek2_dense_and_forward() {
441        let h = 16usize;
442        let n_heads = 2usize;
443        let q_lora = 8usize;
444        let kv_lora = 4usize;
445        let qk_nope = 4usize;
446        let qk_rope = 2usize;
447        let v_dim = 4usize;
448        let ffn = 32usize;
449        let vocab = 8usize;
450        let q_head = qk_nope + qk_rope;
451        let arch = "deepseek2";
452
453        let mut tensors = vec![
454            f32_tensor(
455                "token_embd.weight",
456                vec![vocab as u64, h as u64],
457                vec![0.01; h * vocab],
458            ),
459            f32_tensor("output_norm.weight", vec![h as u64], vec![1.0; h]),
460            f32_tensor(
461                "output.weight",
462                vec![vocab as u64, h as u64],
463                vec![0.02; h * vocab],
464            ),
465        ];
466        for l in 0..2usize {
467            tensors.push(f32_tensor(
468                &format!("blk.{l}.attn_norm.weight"),
469                vec![h as u64],
470                vec![1.0; h],
471            ));
472            tensors.push(f32_tensor(
473                &format!("blk.{l}.ffn_norm.weight"),
474                vec![h as u64],
475                vec![1.0; h],
476            ));
477            tensors.push(f32_tensor(
478                &format!("blk.{l}.attn_q_a.weight"),
479                vec![q_lora as u64, h as u64],
480                vec![0.01; h * q_lora],
481            ));
482            tensors.push(f32_tensor(
483                &format!("blk.{l}.attn_q_b.weight"),
484                vec![(n_heads * q_head) as u64, q_lora as u64],
485                vec![0.01; q_lora * n_heads * q_head],
486            ));
487            tensors.push(f32_tensor(
488                &format!("blk.{l}.attn_kv_a_mqa.weight"),
489                vec![(kv_lora + qk_rope) as u64, h as u64],
490                vec![0.01; h * (kv_lora + qk_rope)],
491            ));
492            tensors.push(f32_tensor(
493                &format!("blk.{l}.attn_kv_b.weight"),
494                vec![(n_heads * (qk_nope + v_dim)) as u64, kv_lora as u64],
495                vec![0.01; kv_lora * n_heads * (qk_nope + v_dim)],
496            ));
497            tensors.push(f32_tensor(
498                &format!("blk.{l}.attn_output.weight"),
499                vec![h as u64, (n_heads * v_dim) as u64],
500                vec![0.01; n_heads * v_dim * h],
501            ));
502            tensors.push(f32_tensor(
503                &format!("blk.{l}.ffn_gate.weight"),
504                vec![ffn as u64, h as u64],
505                vec![0.01; h * ffn],
506            ));
507            tensors.push(f32_tensor(
508                &format!("blk.{l}.ffn_up.weight"),
509                vec![ffn as u64, h as u64],
510                vec![0.01; h * ffn],
511            ));
512            tensors.push(f32_tensor(
513                &format!("blk.{l}.ffn_down.weight"),
514                vec![h as u64, ffn as u64],
515                vec![0.01; ffn * h],
516            ));
517        }
518
519        let kv = [
520            ("deepseek2.block_count", 2u64),
521            ("deepseek2.embedding_length", h as u64),
522            ("deepseek2.feed_forward_length", ffn as u64),
523            ("deepseek2.attention.head_count", n_heads as u64),
524            ("deepseek2.attention.q_lora_rank", q_lora as u64),
525            ("deepseek2.attention.kv_lora_rank", kv_lora as u64),
526            ("deepseek2.attention.qk_nope_head_dim", qk_nope as u64),
527            ("deepseek2.attention.qk_rope_head_dim", qk_rope as u64),
528            ("deepseek2.attention.v_head_dim", v_dim as u64),
529            ("deepseek2.leading_dense_block_count", 2u64),
530            ("deepseek2.expert_count", 0u64),
531        ];
532        let fkv = [
533            ("deepseek2.attention.layer_norm_rms_epsilon", 1e-5f32),
534            ("deepseek2.rope.freq_base", 10000.0f32),
535        ];
536        let bytes = build_gguf(arch, &kv, &fkv, &tensors);
537        let path =
538            std::env::temp_dir().join(format!("ferrox_mla_gguf_{}.gguf", std::process::id()));
539        std::fs::write(&path, &bytes).unwrap();
540        let file = GgufFile::open(&path).unwrap();
541        let engine = load_mla_engine(&file).expect("load mla");
542        assert_eq!(engine.layers.len(), 2);
543        assert_eq!(engine.vocab_size(), vocab);
544        let mut state = engine.new_state();
545        let logits = engine.forward_token(0, 0, &mut state);
546        assert_eq!(logits.len(), vocab);
547        assert!(logits.iter().all(|x| x.is_finite()));
548        let _ = std::fs::remove_file(&path);
549    }
550
551    #[allow(clippy::too_many_arguments)] // test fixture: mirrors the MLA tensor shape set
552    fn push_mla_attn_tensors(
553        tensors: &mut Vec<FixtureTensor>,
554        l: usize,
555        h: usize,
556        n_heads: usize,
557        q_lora: usize,
558        kv_lora: usize,
559        qk_nope: usize,
560        qk_rope: usize,
561        v_dim: usize,
562    ) {
563        let q_head = qk_nope + qk_rope;
564        tensors.push(f32_tensor(
565            &format!("blk.{l}.attn_norm.weight"),
566            vec![h as u64],
567            vec![1.0; h],
568        ));
569        tensors.push(f32_tensor(
570            &format!("blk.{l}.ffn_norm.weight"),
571            vec![h as u64],
572            vec![1.0; h],
573        ));
574        tensors.push(f32_tensor(
575            &format!("blk.{l}.attn_q_a.weight"),
576            vec![q_lora as u64, h as u64],
577            vec![0.01; h * q_lora],
578        ));
579        tensors.push(f32_tensor(
580            &format!("blk.{l}.attn_q_b.weight"),
581            vec![(n_heads * q_head) as u64, q_lora as u64],
582            vec![0.01; q_lora * n_heads * q_head],
583        ));
584        tensors.push(f32_tensor(
585            &format!("blk.{l}.attn_kv_a_mqa.weight"),
586            vec![(kv_lora + qk_rope) as u64, h as u64],
587            vec![0.01; h * (kv_lora + qk_rope)],
588        ));
589        tensors.push(f32_tensor(
590            &format!("blk.{l}.attn_kv_b.weight"),
591            vec![(n_heads * (qk_nope + v_dim)) as u64, kv_lora as u64],
592            vec![0.01; kv_lora * n_heads * (qk_nope + v_dim)],
593        ));
594        tensors.push(f32_tensor(
595            &format!("blk.{l}.attn_output.weight"),
596            vec![h as u64, (n_heads * v_dim) as u64],
597            vec![0.01; n_heads * v_dim * h],
598        ));
599    }
600
601    #[test]
602    fn load_synthetic_deepseek2_moe_after_dense_and_forward() {
603        let h = 16usize;
604        let n_heads = 2usize;
605        let q_lora = 8usize;
606        let kv_lora = 4usize;
607        let qk_nope = 4usize;
608        let qk_rope = 2usize;
609        let v_dim = 4usize;
610        let ffn = 32usize;
611        let exp_ff = 16usize;
612        let n_exp = 4usize;
613        let vocab = 8usize;
614        let arch = "deepseek2";
615
616        let mut tensors = vec![
617            f32_tensor(
618                "token_embd.weight",
619                vec![vocab as u64, h as u64],
620                vec![0.01; h * vocab],
621            ),
622            f32_tensor("output_norm.weight", vec![h as u64], vec![1.0; h]),
623            f32_tensor(
624                "output.weight",
625                vec![vocab as u64, h as u64],
626                vec![0.02; h * vocab],
627            ),
628        ];
629        // Layer 0: dense
630        push_mla_attn_tensors(
631            &mut tensors,
632            0,
633            h,
634            n_heads,
635            q_lora,
636            kv_lora,
637            qk_nope,
638            qk_rope,
639            v_dim,
640        );
641        tensors.push(f32_tensor(
642            "blk.0.ffn_gate.weight",
643            vec![ffn as u64, h as u64],
644            vec![0.01; h * ffn],
645        ));
646        tensors.push(f32_tensor(
647            "blk.0.ffn_up.weight",
648            vec![ffn as u64, h as u64],
649            vec![0.01; h * ffn],
650        ));
651        tensors.push(f32_tensor(
652            "blk.0.ffn_down.weight",
653            vec![h as u64, ffn as u64],
654            vec![0.01; ffn * h],
655        ));
656        // Layer 1: MoE
657        push_mla_attn_tensors(
658            &mut tensors,
659            1,
660            h,
661            n_heads,
662            q_lora,
663            kv_lora,
664            qk_nope,
665            qk_rope,
666            v_dim,
667        );
668        tensors.push(f32_tensor(
669            "blk.1.ffn_gate_inp.weight",
670            vec![n_exp as u64, h as u64],
671            vec![0.01; h * n_exp],
672        ));
673        // Packed expert tensors: logical [n_experts, out, in] → GGUF shape write uses rev
674        // so pass shape as [n_experts, out, in] matching other fixtures' logical order.
675        tensors.push(f32_tensor(
676            "blk.1.ffn_gate_exps.weight",
677            vec![n_exp as u64, exp_ff as u64, h as u64],
678            vec![0.01; n_exp * exp_ff * h],
679        ));
680        tensors.push(f32_tensor(
681            "blk.1.ffn_up_exps.weight",
682            vec![n_exp as u64, exp_ff as u64, h as u64],
683            vec![0.01; n_exp * exp_ff * h],
684        ));
685        tensors.push(f32_tensor(
686            "blk.1.ffn_down_exps.weight",
687            vec![n_exp as u64, h as u64, exp_ff as u64],
688            vec![0.01; n_exp * h * exp_ff],
689        ));
690        tensors.push(f32_tensor(
691            "blk.1.ffn_gate_shexp.weight",
692            vec![exp_ff as u64, h as u64],
693            vec![0.01; h * exp_ff],
694        ));
695        tensors.push(f32_tensor(
696            "blk.1.ffn_up_shexp.weight",
697            vec![exp_ff as u64, h as u64],
698            vec![0.01; h * exp_ff],
699        ));
700        tensors.push(f32_tensor(
701            "blk.1.ffn_down_shexp.weight",
702            vec![h as u64, exp_ff as u64],
703            vec![0.01; exp_ff * h],
704        ));
705
706        let kv = [
707            ("deepseek2.block_count", 2u64),
708            ("deepseek2.embedding_length", h as u64),
709            ("deepseek2.feed_forward_length", ffn as u64),
710            ("deepseek2.attention.head_count", n_heads as u64),
711            ("deepseek2.attention.q_lora_rank", q_lora as u64),
712            ("deepseek2.attention.kv_lora_rank", kv_lora as u64),
713            ("deepseek2.attention.qk_nope_head_dim", qk_nope as u64),
714            ("deepseek2.attention.qk_rope_head_dim", qk_rope as u64),
715            ("deepseek2.attention.v_head_dim", v_dim as u64),
716            ("deepseek2.leading_dense_block_count", 1u64),
717            ("deepseek2.expert_count", n_exp as u64),
718            ("deepseek2.expert_used_count", 2u64),
719            ("deepseek2.expert_shared_count", 1u64),
720            ("deepseek2.expert_feed_forward_length", exp_ff as u64),
721            ("deepseek2.expert_gating_func", 1u64), // softmax
722        ];
723        let fkv = [
724            ("deepseek2.attention.layer_norm_rms_epsilon", 1e-5f32),
725            ("deepseek2.rope.freq_base", 10000.0f32),
726            ("deepseek2.expert_weights_scale", 1.0f32),
727        ];
728        let bytes = build_gguf(arch, &kv, &fkv, &tensors);
729        let path =
730            std::env::temp_dir().join(format!("ferrox_mla_moe_gguf_{}.gguf", std::process::id()));
731        std::fs::write(&path, &bytes).unwrap();
732        let file = GgufFile::open(&path).unwrap();
733        let engine = load_mla_engine(&file).expect("load mla moe");
734        assert_eq!(engine.layers.len(), 2);
735        assert!(matches!(
736            engine.layers[0].ffn,
737            crate::engine::MlaLayerFfn::Dense(_)
738        ));
739        assert!(matches!(
740            engine.layers[1].ffn,
741            crate::engine::MlaLayerFfn::Moe(_)
742        ));
743        assert!(engine.moe.is_some());
744        let mut state = engine.new_state();
745        let logits = engine.forward_token(0, 0, &mut state);
746        assert_eq!(logits.len(), vocab);
747        assert!(logits.iter().all(|x| x.is_finite()));
748        let _ = std::fs::remove_file(&path);
749    }
750
751    #[test]
752    fn moe_after_dense_fails_closed_without_expert_tensors() {
753        let h = 16usize;
754        let n_heads = 2usize;
755        let q_lora = 8usize;
756        let kv_lora = 4usize;
757        let qk_nope = 4usize;
758        let qk_rope = 2usize;
759        let v_dim = 4usize;
760        let ffn = 32usize;
761        let vocab = 8usize;
762        let arch = "deepseek2";
763
764        let mut tensors = vec![
765            f32_tensor(
766                "token_embd.weight",
767                vec![vocab as u64, h as u64],
768                vec![0.01; h * vocab],
769            ),
770            f32_tensor("output_norm.weight", vec![h as u64], vec![1.0; h]),
771            f32_tensor(
772                "output.weight",
773                vec![vocab as u64, h as u64],
774                vec![0.02; h * vocab],
775            ),
776        ];
777        for l in 0..2usize {
778            push_mla_attn_tensors(
779                &mut tensors,
780                l,
781                h,
782                n_heads,
783                q_lora,
784                kv_lora,
785                qk_nope,
786                qk_rope,
787                v_dim,
788            );
789            // Only dense FFN tensors — MoE layer 1 will fail closed.
790            tensors.push(f32_tensor(
791                &format!("blk.{l}.ffn_gate.weight"),
792                vec![ffn as u64, h as u64],
793                vec![0.01; h * ffn],
794            ));
795            tensors.push(f32_tensor(
796                &format!("blk.{l}.ffn_up.weight"),
797                vec![ffn as u64, h as u64],
798                vec![0.01; h * ffn],
799            ));
800            tensors.push(f32_tensor(
801                &format!("blk.{l}.ffn_down.weight"),
802                vec![h as u64, ffn as u64],
803                vec![0.01; ffn * h],
804            ));
805        }
806        let kv = [
807            ("deepseek2.block_count", 2u64),
808            ("deepseek2.embedding_length", h as u64),
809            ("deepseek2.feed_forward_length", ffn as u64),
810            ("deepseek2.attention.head_count", n_heads as u64),
811            ("deepseek2.attention.q_lora_rank", q_lora as u64),
812            ("deepseek2.attention.kv_lora_rank", kv_lora as u64),
813            ("deepseek2.attention.qk_nope_head_dim", qk_nope as u64),
814            ("deepseek2.attention.qk_rope_head_dim", qk_rope as u64),
815            ("deepseek2.attention.v_head_dim", v_dim as u64),
816            ("deepseek2.leading_dense_block_count", 1u64),
817            ("deepseek2.expert_count", 4u64),
818            ("deepseek2.expert_used_count", 2u64),
819        ];
820        let fkv = [
821            ("deepseek2.attention.layer_norm_rms_epsilon", 1e-5f32),
822            ("deepseek2.rope.freq_base", 10000.0f32),
823        ];
824        let bytes = build_gguf(arch, &kv, &fkv, &tensors);
825        let path = std::env::temp_dir().join(format!(
826            "ferrox_mla_moe_missing_{}.gguf",
827            std::process::id()
828        ));
829        std::fs::write(&path, &bytes).unwrap();
830        let file = GgufFile::open(&path).unwrap();
831        let err = match load_mla_engine(&file) {
832            Err(e) => e,
833            Ok(_) => panic!("expected missing MoE tensors to fail closed"),
834        };
835        let msg = format!("{err}");
836        assert!(
837            msg.contains("ffn_gate_inp") || msg.contains("TensorNotFound"),
838            "unexpected error: {msg}"
839        );
840        let _ = std::fs::remove_file(&path);
841    }
842}