1use 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
62fn 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]; let mut transposed = vec![0f32; per_head]; 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
111fn 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
171pub 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
193pub 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
329pub 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
350pub 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 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
405pub 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
417pub 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
526pub 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 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(); buf.write_u64::<LittleEndian>(tensors.len() as u64).unwrap();
624 buf.write_u64::<LittleEndian>(1).unwrap(); 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(); 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(); 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 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 let tensors = vec![f32_tensor(
930 "blk.0.attn_k_b.weight",
931 vec![1, 3, 2], 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 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}