1use 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#[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 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
59pub 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 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 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 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 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
277pub 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 } 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 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(); 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(); 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)] 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 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 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 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), ];
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 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}