1use ferrox_core::tensor::Tensor;
43use ferrox_core::weight_matrix::{WeightBytes, WeightMatrix};
44use ferrox_safetensors::{SafetensorsDtype, ShardedSafetensors};
45use thiserror::Error;
46
47use crate::config::LayerAttentionKind;
48use crate::kda::KdaAttnWeights;
49use crate::kimi_decoder::DenseMlpWeights;
50use crate::latent_moe::{KimiExpertBacking, KimiExpertWeights, KimiLatentMoeWeights};
51use crate::mla::MlaAttnWeights;
52use ferrox_core::expert_store::{ExpertKey, ExpertSource, ExpertStore};
53
54#[derive(Debug, Error)]
55pub enum KimiLoadError {
56 #[error("safetensors error: {0}")]
57 Safetensors(#[from] ferrox_safetensors::SafetensorsError),
58 #[error("tensor '{0}' has unsupported dtype {1:?} (expected F32 or BF16)")]
59 UnsupportedDtype(String, SafetensorsDtype),
60 #[error("{0}")]
61 Other(String),
62}
63
64pub fn load_f32_vec(shard: &ShardedSafetensors, name: &str) -> Result<Vec<f32>, KimiLoadError> {
72 let info = shard
73 .tensor_info(name)
74 .ok_or_else(|| ferrox_safetensors::SafetensorsError::TensorNotFound(name.to_string()))?;
75 let raw = shard.tensor_bytes(name)?;
76 match info.dtype {
77 SafetensorsDtype::F32 => {
78 let mut out = Vec::with_capacity(raw.len() / 4);
79 for chunk in raw.chunks_exact(4) {
80 out.push(f32::from_le_bytes(chunk.try_into().unwrap()));
81 }
82 Ok(out)
83 }
84 SafetensorsDtype::BF16 => ferrox_quant::dequant_bf16(raw)
85 .map_err(|_| KimiLoadError::UnsupportedDtype(name.to_string(), info.dtype)),
86 other => Err(KimiLoadError::UnsupportedDtype(name.to_string(), other)),
87 }
88}
89
90fn load_weight_matrix(
91 shard: &ShardedSafetensors,
92 name: &str,
93 rows: usize,
94 cols: usize,
95) -> Result<WeightMatrix, KimiLoadError> {
96 let data = load_f32_vec(shard, name)?;
97 assert_eq!(
98 data.len(),
99 rows * cols,
100 "tensor '{name}' has {} elements, expected {rows}*{cols}",
101 data.len()
102 );
103 Ok(WeightMatrix::F32(Tensor::new(data, vec![rows, cols])))
104}
105
106pub fn load_kda_attn(
112 shard: &ShardedSafetensors,
113 prefix: &str,
114 num_heads: usize,
115 head_dim: usize,
116 hidden_dim: usize,
117) -> Result<KdaAttnWeights, KimiLoadError> {
118 let projection_size = num_heads * head_dim;
119 let a_log_full = load_f32_vec(shard, &format!("{prefix}.self_attn.A_log"))?;
120 let a_log = a_log_full[..num_heads].to_vec();
124
125 Ok(KdaAttnWeights {
126 q_proj: load_weight_matrix(
127 shard,
128 &format!("{prefix}.self_attn.q_proj.weight"),
129 projection_size,
130 hidden_dim,
131 )?,
132 k_proj: load_weight_matrix(
133 shard,
134 &format!("{prefix}.self_attn.k_proj.weight"),
135 projection_size,
136 hidden_dim,
137 )?,
138 v_proj: load_weight_matrix(
139 shard,
140 &format!("{prefix}.self_attn.v_proj.weight"),
141 projection_size,
142 hidden_dim,
143 )?,
144 q_conv_weight: load_f32_vec(shard, &format!("{prefix}.self_attn.q_conv1d.weight"))?,
149 k_conv_weight: load_f32_vec(shard, &format!("{prefix}.self_attn.k_conv1d.weight"))?,
150 v_conv_weight: load_f32_vec(shard, &format!("{prefix}.self_attn.v_conv1d.weight"))?,
151 a_log,
152 f_a_proj: load_weight_matrix(
153 shard,
154 &format!("{prefix}.self_attn.f_a_proj.weight"),
155 head_dim,
156 hidden_dim,
157 )?,
158 f_b_proj: load_weight_matrix(
159 shard,
160 &format!("{prefix}.self_attn.f_b_proj.weight"),
161 projection_size,
162 head_dim,
163 )?,
164 dt_bias: load_f32_vec(shard, &format!("{prefix}.self_attn.dt_bias"))?,
165 b_proj: load_weight_matrix(
166 shard,
167 &format!("{prefix}.self_attn.b_proj.weight"),
168 num_heads,
169 hidden_dim,
170 )?,
171 g_proj: load_weight_matrix(
172 shard,
173 &format!("{prefix}.self_attn.g_proj.weight"),
174 projection_size,
175 hidden_dim,
176 )?,
177 o_norm_weight: load_f32_vec(shard, &format!("{prefix}.self_attn.o_norm.weight"))?,
178 o_proj: load_weight_matrix(
179 shard,
180 &format!("{prefix}.self_attn.o_proj.weight"),
181 hidden_dim,
182 projection_size,
183 )?,
184 })
185}
186
187#[allow(clippy::too_many_arguments)]
190pub fn load_mla_attn(
191 shard: &ShardedSafetensors,
192 prefix: &str,
193 num_heads: usize,
194 q_lora_rank: usize,
195 kv_lora_rank: usize,
196 qk_nope_head_dim: usize,
197 qk_rope_head_dim: usize,
198 v_head_dim: usize,
199 hidden_dim: usize,
200) -> Result<MlaAttnWeights, KimiLoadError> {
201 let q_head_dim = qk_nope_head_dim + qk_rope_head_dim;
202 Ok(MlaAttnWeights {
203 q_a_proj: load_weight_matrix(
204 shard,
205 &format!("{prefix}.self_attn.q_a_proj.weight"),
206 q_lora_rank,
207 hidden_dim,
208 )?,
209 q_a_layernorm: load_f32_vec(shard, &format!("{prefix}.self_attn.q_a_layernorm.weight"))?,
210 q_b_proj: load_weight_matrix(
211 shard,
212 &format!("{prefix}.self_attn.q_b_proj.weight"),
213 num_heads * q_head_dim,
214 q_lora_rank,
215 )?,
216 kv_a_proj_with_mqa: load_weight_matrix(
217 shard,
218 &format!("{prefix}.self_attn.kv_a_proj_with_mqa.weight"),
219 kv_lora_rank + qk_rope_head_dim,
220 hidden_dim,
221 )?,
222 kv_a_layernorm: load_f32_vec(shard, &format!("{prefix}.self_attn.kv_a_layernorm.weight"))?,
223 kv_b_proj: load_weight_matrix(
224 shard,
225 &format!("{prefix}.self_attn.kv_b_proj.weight"),
226 num_heads * (qk_nope_head_dim + v_head_dim),
227 kv_lora_rank,
228 )?,
229 o_proj: load_weight_matrix(
230 shard,
231 &format!("{prefix}.self_attn.o_proj.weight"),
232 hidden_dim,
233 num_heads * v_head_dim,
234 )?,
235 g_proj: Some(load_weight_matrix(
236 shard,
237 &format!("{prefix}.self_attn.g_proj.weight"),
238 num_heads * v_head_dim,
239 hidden_dim,
240 )?),
241 })
242}
243
244pub fn load_dense_mlp(
248 shard: &ShardedSafetensors,
249 prefix: &str,
250 hidden_dim: usize,
251 intermediate_dim: usize,
252) -> Result<DenseMlpWeights, KimiLoadError> {
253 Ok(DenseMlpWeights {
254 gate_proj: load_weight_matrix(
255 shard,
256 &format!("{prefix}.mlp.gate_proj.weight"),
257 intermediate_dim,
258 hidden_dim,
259 )?,
260 up_proj: load_weight_matrix(
261 shard,
262 &format!("{prefix}.mlp.up_proj.weight"),
263 intermediate_dim,
264 hidden_dim,
265 )?,
266 down_proj: load_weight_matrix(
267 shard,
268 &format!("{prefix}.mlp.down_proj.weight"),
269 hidden_dim,
270 intermediate_dim,
271 )?,
272 })
273}
274
275pub struct BlockResidualWeights {
281 pub self_attention_res_norm_weight: Vec<f32>,
282 pub self_attention_res_proj_weight: Vec<f32>,
283 pub mlp_res_norm_weight: Vec<f32>,
284 pub mlp_res_proj_weight: Vec<f32>,
285}
286
287pub fn load_block_residual(
288 shard: &ShardedSafetensors,
289 prefix: &str,
290) -> Result<BlockResidualWeights, KimiLoadError> {
291 Ok(BlockResidualWeights {
292 self_attention_res_norm_weight: load_f32_vec(
293 shard,
294 &format!("{prefix}.self_attention_res_norm.weight"),
295 )?,
296 self_attention_res_proj_weight: load_f32_vec(
297 shard,
298 &format!("{prefix}.self_attention_res_proj.weight"),
299 )?,
300 mlp_res_norm_weight: load_f32_vec(shard, &format!("{prefix}.mlp_res_norm.weight"))?,
301 mlp_res_proj_weight: load_f32_vec(shard, &format!("{prefix}.mlp_res_proj.weight"))?,
302 })
303}
304
305fn load_mxfp4_weight_matrix(
316 shard: &ShardedSafetensors,
317 packed_name: &str,
318 scale_name: &str,
319 rows: usize,
320 cols: usize,
321) -> Result<WeightMatrix, KimiLoadError> {
322 let (packed_mmap, packed_range) = shard.tensor_mapped_range(packed_name)?;
323 let (scale_mmap, scale_range) = shard.tensor_mapped_range(scale_name)?;
324 let packed_per_row = cols / 2;
325 let scale_per_row = cols / ferrox_quant::MXFP4_GROUP_SIZE;
326 assert_eq!(
327 packed_range.len(),
328 rows * packed_per_row,
329 "'{packed_name}' has {} bytes, expected {rows}*{packed_per_row}",
330 packed_range.len()
331 );
332 assert_eq!(
333 scale_range.len(),
334 rows * scale_per_row,
335 "'{scale_name}' has {} bytes, expected {rows}*{scale_per_row}",
336 scale_range.len()
337 );
338
339 Ok(WeightMatrix::Mxfp4 {
340 packed: WeightBytes::Mapped {
341 mmap: packed_mmap,
342 range: packed_range,
343 },
344 scale: WeightBytes::Mapped {
345 mmap: scale_mmap,
346 range: scale_range,
347 },
348 rows,
349 cols,
350 })
351}
352
353#[derive(Debug, Clone, Copy)]
364pub struct KimiStoredExpertLayout {
365 pub moe_hidden_dim: usize,
366 pub moe_intermediate_dim: usize,
367}
368
369impl KimiStoredExpertLayout {
370 fn seg_lens(&self) -> [usize; 6] {
371 let (h, m) = (self.moe_hidden_dim, self.moe_intermediate_dim);
372 let g = ferrox_quant::MXFP4_GROUP_SIZE;
373 [
374 m * h / 2, m * h / g, h * m / 2, h * m / g, m * h / 2, m * h / g, ]
381 }
382
383 pub fn total_bytes(&self) -> usize {
384 self.seg_lens().iter().sum()
385 }
386
387 pub fn materialize(&self, lease: &ferrox_core::expert_store::ExpertLease) -> KimiExpertWeights {
391 let (h, m) = (self.moe_hidden_dim, self.moe_intermediate_dim);
392 let lens = self.seg_lens();
393 let mut offsets = [0usize; 6];
394 for i in 1..6 {
395 offsets[i] = offsets[i - 1] + lens[i - 1];
396 }
397 let shared = |i: usize| WeightBytes::Shared {
398 buf: lease.shared_buf(),
399 range: offsets[i]..offsets[i] + lens[i],
400 };
401 let mx = |pi: usize, si: usize, rows: usize, cols: usize| WeightMatrix::Mxfp4 {
402 packed: shared(pi),
403 scale: shared(si),
404 rows,
405 cols,
406 };
407 KimiExpertWeights {
408 w1: mx(0, 1, m, h),
409 w2: mx(2, 3, h, m),
410 w3: mx(4, 5, m, h),
411 }
412 }
413}
414
415pub struct KimiExpertSource {
420 files: Vec<std::fs::File>,
421 segments: std::collections::HashMap<ExpertKey, [(usize, u64, usize); 6]>,
423}
424
425impl ExpertSource for KimiExpertSource {
426 fn expert_len(&self, key: ExpertKey) -> Option<usize> {
427 self.segments
428 .get(&key)
429 .map(|segs| segs.iter().map(|&(_, _, len)| len).sum())
430 }
431
432 fn read_expert(&self, key: ExpertKey) -> std::io::Result<Vec<u8>> {
433 let segs = self
434 .segments
435 .get(&key)
436 .ok_or_else(|| std::io::Error::new(std::io::ErrorKind::NotFound, format!("{key:?}")))?;
437 let total: usize = segs.iter().map(|&(_, _, len)| len).sum();
438 let mut buf = vec![0u8; total];
439 let mut written = 0;
440 for &(fi, offset, len) in segs {
441 let dst = &mut buf[written..written + len];
442 #[cfg(unix)]
443 {
444 use std::os::unix::fs::FileExt;
445 self.files[fi].read_exact_at(dst, offset)?;
446 }
447 #[cfg(not(unix))]
448 {
449 use std::io::{Read, Seek, SeekFrom};
450 let mut f = &self.files[fi];
451 f.seek(SeekFrom::Start(offset))?;
452 f.read_exact(dst)?;
453 }
454 written += len;
455 }
456 Ok(buf)
457 }
458}
459
460pub fn load_kimi_expert(
461 shard: &ShardedSafetensors,
462 moe_prefix: &str,
463 expert_idx: usize,
464 moe_hidden_dim: usize,
465 moe_intermediate_dim: usize,
466) -> Result<KimiExpertWeights, KimiLoadError> {
467 let expert_prefix = format!("{moe_prefix}.experts.{expert_idx}");
468 Ok(KimiExpertWeights {
469 w1: load_mxfp4_weight_matrix(
470 shard,
471 &format!("{expert_prefix}.w1.weight_packed"),
472 &format!("{expert_prefix}.w1.weight_scale"),
473 moe_intermediate_dim,
474 moe_hidden_dim,
475 )?,
476 w2: load_mxfp4_weight_matrix(
477 shard,
478 &format!("{expert_prefix}.w2.weight_packed"),
479 &format!("{expert_prefix}.w2.weight_scale"),
480 moe_hidden_dim,
481 moe_intermediate_dim,
482 )?,
483 w3: load_mxfp4_weight_matrix(
484 shard,
485 &format!("{expert_prefix}.w3.weight_packed"),
486 &format!("{expert_prefix}.w3.weight_scale"),
487 moe_intermediate_dim,
488 moe_hidden_dim,
489 )?,
490 })
491}
492
493#[allow(clippy::too_many_arguments)]
499pub fn load_latent_moe(
500 shard: &ShardedSafetensors,
501 prefix: &str,
502 hidden_dim: usize,
503 moe_hidden_dim: usize,
504 moe_intermediate_dim: usize,
505 n_experts: usize,
506 shared_intermediate_dim: usize,
507) -> Result<KimiLatentMoeWeights, KimiLoadError> {
508 let moe_prefix = format!("{prefix}.block_sparse_moe");
509 let mut experts = Vec::with_capacity(n_experts);
510 for e in 0..n_experts {
511 experts.push(load_kimi_expert(
512 shard,
513 &moe_prefix,
514 e,
515 moe_hidden_dim,
516 moe_intermediate_dim,
517 )?);
518 }
519 let experts = KimiExpertBacking::Resident(experts);
520
521 Ok(KimiLatentMoeWeights {
522 router_weight: load_weight_matrix(
523 shard,
524 &format!("{moe_prefix}.gate.weight"),
525 n_experts,
526 hidden_dim,
527 )?,
528 e_score_correction_bias: load_f32_vec(
529 shard,
530 &format!("{moe_prefix}.gate.e_score_correction_bias"),
531 )?,
532 down_proj: load_weight_matrix(
533 shard,
534 &format!("{moe_prefix}.routed_expert_down_proj.weight"),
535 moe_hidden_dim,
536 hidden_dim,
537 )?,
538 up_proj: load_weight_matrix(
539 shard,
540 &format!("{moe_prefix}.routed_expert_up_proj.weight"),
541 hidden_dim,
542 moe_hidden_dim,
543 )?,
544 routed_expert_norm_weight: Some(load_f32_vec(
545 shard,
546 &format!("{moe_prefix}.routed_expert_norm.weight"),
547 )?),
548 experts,
549 shared_expert: KimiExpertWeights {
550 w1: load_weight_matrix(
551 shard,
552 &format!("{moe_prefix}.shared_experts.gate_proj.weight"),
553 shared_intermediate_dim,
554 hidden_dim,
555 )?,
556 w2: load_weight_matrix(
557 shard,
558 &format!("{moe_prefix}.shared_experts.down_proj.weight"),
559 hidden_dim,
560 shared_intermediate_dim,
561 )?,
562 w3: load_weight_matrix(
563 shard,
564 &format!("{moe_prefix}.shared_experts.up_proj.weight"),
565 shared_intermediate_dim,
566 hidden_dim,
567 )?,
568 },
569 })
570}
571
572pub struct KimiRealHparams {
579 pub hidden_dim: usize,
580 pub kda_num_heads: usize,
581 pub kda_head_dim: usize,
582 pub mla_num_heads: usize,
583 pub mla_q_lora_rank: usize,
584 pub mla_kv_lora_rank: usize,
585 pub mla_qk_nope_head_dim: usize,
586 pub mla_qk_rope_head_dim: usize,
587 pub mla_v_head_dim: usize,
588 pub dense_intermediate_dim: usize,
589 pub moe_hidden_dim: usize,
590 pub moe_intermediate_dim: usize,
591 pub n_experts: usize,
592 pub num_shared_experts: usize,
593}
594
595impl KimiRealHparams {
596 pub fn real() -> Self {
598 KimiRealHparams {
599 hidden_dim: 7168,
600 kda_num_heads: 96,
601 kda_head_dim: 128,
602 mla_num_heads: 96,
603 mla_q_lora_rank: 1536,
604 mla_kv_lora_rank: 512,
605 mla_qk_nope_head_dim: 128,
606 mla_qk_rope_head_dim: 64,
607 mla_v_head_dim: 128,
608 dense_intermediate_dim: 33792,
609 moe_hidden_dim: 3584,
610 moe_intermediate_dim: 3072,
611 n_experts: 896,
612 num_shared_experts: 2,
613 }
614 }
615}
616
617pub fn load_kimi_layer(
623 shard: &ShardedSafetensors,
624 hp: &KimiRealHparams,
625 kind: LayerAttentionKind,
626 is_dense: bool,
627 layer_idx: usize,
628) -> Result<crate::kimi_decoder::KimiDecoderLayerWeights, KimiLoadError> {
629 let prefix = format!("language_model.model.layers.{layer_idx}");
630
631 let input_layernorm_weight = load_f32_vec(shard, &format!("{prefix}.input_layernorm.weight"))?;
632 let post_attention_layernorm_weight =
633 load_f32_vec(shard, &format!("{prefix}.post_attention_layernorm.weight"))?;
634 let block_res = load_block_residual(shard, &prefix)?;
635
636 let attn = match kind {
637 LayerAttentionKind::KimiKda => {
638 crate::kimi_decoder::KimiLayerAttention::Kda(Box::new(load_kda_attn(
639 shard,
640 &prefix,
641 hp.kda_num_heads,
642 hp.kda_head_dim,
643 hp.hidden_dim,
644 )?))
645 }
646 LayerAttentionKind::KimiMla => {
647 crate::kimi_decoder::KimiLayerAttention::Mla(Box::new(load_mla_attn(
648 shard,
649 &prefix,
650 hp.mla_num_heads,
651 hp.mla_q_lora_rank,
652 hp.mla_kv_lora_rank,
653 hp.mla_qk_nope_head_dim,
654 hp.mla_qk_rope_head_dim,
655 hp.mla_v_head_dim,
656 hp.hidden_dim,
657 )?))
658 }
659 LayerAttentionKind::Gqa => {
660 panic!("load_kimi_layer is only for KimiHybrid (KDA/Gated-MLA) layers")
661 }
662 };
663
664 let ffn = if is_dense {
665 crate::kimi_decoder::KimiLayerFfn::Dense(Box::new(load_dense_mlp(
666 shard,
667 &prefix,
668 hp.hidden_dim,
669 hp.dense_intermediate_dim,
670 )?))
671 } else {
672 crate::kimi_decoder::KimiLayerFfn::Moe(Box::new(load_latent_moe(
673 shard,
674 &prefix,
675 hp.hidden_dim,
676 hp.moe_hidden_dim,
677 hp.moe_intermediate_dim,
678 hp.n_experts,
679 hp.moe_intermediate_dim * hp.num_shared_experts,
680 )?))
681 };
682
683 Ok(crate::kimi_decoder::KimiDecoderLayerWeights {
684 input_layernorm_weight,
685 attn,
686 post_attention_layernorm_weight,
687 ffn,
688 self_attention_res_norm_weight: block_res.self_attention_res_norm_weight,
689 self_attention_res_proj_weight: block_res.self_attention_res_proj_weight,
690 mlp_res_norm_weight: block_res.mlp_res_norm_weight,
691 mlp_res_proj_weight: block_res.mlp_res_proj_weight,
692 })
693}
694
695pub fn load_kimi_checkpoint_with_expert_cache(
723 shard: &ShardedSafetensors,
724 model_cfg: &crate::config::ModelConfig,
725 hp: &KimiRealHparams,
726 expert_cache_bytes: Option<u64>,
727) -> Result<crate::kimi_decoder::KimiDecoderWeights, KimiLoadError> {
728 let mut weights = load_kimi_checkpoint(shard, model_cfg, hp)?;
729 let Some(budget) = expert_cache_bytes else {
730 return Ok(weights);
731 };
732
733 let mut files: Vec<std::fs::File> = Vec::new();
736 let mut path_index: std::collections::HashMap<std::path::PathBuf, usize> =
737 std::collections::HashMap::new();
738 let mut segments: std::collections::HashMap<ExpertKey, [(usize, u64, usize); 6]> =
739 std::collections::HashMap::new();
740 let layout = KimiStoredExpertLayout {
741 moe_hidden_dim: hp.moe_hidden_dim,
742 moe_intermediate_dim: hp.moe_intermediate_dim,
743 };
744 let mut moe_layers: Vec<(usize, usize)> = Vec::new(); for (layer_idx, layer) in weights.layers.iter().enumerate() {
747 let crate::kimi_decoder::KimiLayerFfn::Moe(moe) = &layer.ffn else {
748 continue;
749 };
750 let n_experts = moe.experts.n_experts();
751 let moe_prefix = format!("language_model.model.layers.{layer_idx}.block_sparse_moe");
752 for e in 0..n_experts {
753 let expert_prefix = format!("{moe_prefix}.experts.{e}");
754 let mut segs = [(0usize, 0u64, 0usize); 6];
755 for (i, tensor) in [
756 format!("{expert_prefix}.w1.weight_packed"),
757 format!("{expert_prefix}.w1.weight_scale"),
758 format!("{expert_prefix}.w2.weight_packed"),
759 format!("{expert_prefix}.w2.weight_scale"),
760 format!("{expert_prefix}.w3.weight_packed"),
761 format!("{expert_prefix}.w3.weight_scale"),
762 ]
763 .iter()
764 .enumerate()
765 {
766 let (path, range) = shard.tensor_file_location(tensor)?;
767 let fi = match path_index.get(path) {
768 Some(&fi) => fi,
769 None => {
770 let fi = files.len();
771 files.push(std::fs::File::open(path).map_err(|e| {
772 KimiLoadError::Other(format!(
773 "opening shard file {} for expert streaming: {e}",
774 path.display()
775 ))
776 })?);
777 path_index.insert(path.to_path_buf(), fi);
778 fi
779 }
780 };
781 segs[i] = (fi, range.start as u64, range.end - range.start);
782 }
783 segments.insert(
784 ExpertKey {
785 layer: layer_idx as u32,
786 expert: e as u32,
787 },
788 segs,
789 );
790 }
791 moe_layers.push((layer_idx, n_experts));
792 }
793
794 if moe_layers.is_empty() {
795 return Ok(weights);
796 }
797 let store = std::sync::Arc::new(ExpertStore::new(
798 KimiExpertSource { files, segments },
799 budget as usize,
800 ));
801 for (layer_idx, n_experts) in moe_layers {
802 if let crate::kimi_decoder::KimiLayerFfn::Moe(moe) = &mut weights.layers[layer_idx].ffn {
803 moe.experts = KimiExpertBacking::Stored {
804 store: std::sync::Arc::clone(&store),
805 layout,
806 n_experts,
807 layer: layer_idx as u32,
808 };
809 }
810 }
811 Ok(weights)
812}
813
814pub fn load_kimi_checkpoint(
815 shard: &ShardedSafetensors,
816 model_cfg: &crate::config::ModelConfig,
817 hp: &KimiRealHparams,
818) -> Result<crate::kimi_decoder::KimiDecoderWeights, KimiLoadError> {
819 let mut layers = Vec::with_capacity(model_cfg.n_layers);
820 for layer_idx in 0..model_cfg.n_layers {
821 let kind = model_cfg.layer_attention_kind(layer_idx);
822 let is_dense = model_cfg.layer_is_dense(layer_idx);
823 layers.push(load_kimi_layer(shard, hp, kind, is_dense, layer_idx)?);
824 }
825
826 let embedding_data = load_f32_vec(shard, "language_model.model.embed_tokens.weight")?;
827 let embedding = Tensor::new(embedding_data, vec![model_cfg.vocab_size, hp.hidden_dim]);
828 let output_head = load_weight_matrix(
829 shard,
830 "language_model.lm_head.weight",
831 model_cfg.vocab_size,
832 hp.hidden_dim,
833 )?;
834 let final_norm_weight = load_f32_vec(shard, "language_model.model.norm.weight")?;
835 let output_attn_res_norm_weight =
836 load_f32_vec(shard, "language_model.model.output_attn_res_norm.weight")?;
837 let output_attn_res_proj_weight =
838 load_f32_vec(shard, "language_model.model.output_attn_res_proj.weight")?;
839
840 Ok(crate::kimi_decoder::KimiDecoderWeights {
841 embedding,
842 layers,
843 output_attn_res_norm_weight,
844 output_attn_res_proj_weight,
845 final_norm_weight,
846 output_head,
847 })
848}
849
850#[cfg(test)]
851mod tests {
852 use super::*;
853 use byteorder::{LittleEndian, WriteBytesExt};
854 use std::io::Write;
855
856 fn bf16_bytes(values: &[f32]) -> Vec<u8> {
861 let mut out = Vec::with_capacity(values.len() * 2);
862 for &v in values {
863 let bits = v.to_bits();
864 out.extend_from_slice(&((bits >> 16) as u16).to_le_bytes());
865 }
866 out
867 }
868
869 fn build_shard(tensors: &[(&str, &str, &[usize], Vec<u8>)]) -> Vec<u8> {
870 let mut header = String::from("{");
871 let mut offset = 0u64;
872 let mut data = Vec::new();
873 for (i, (name, dtype, shape, bytes)) in tensors.iter().enumerate() {
874 if i > 0 {
875 header.push(',');
876 }
877 let shape_str = shape
878 .iter()
879 .map(|d| d.to_string())
880 .collect::<Vec<_>>()
881 .join(",");
882 let end = offset + bytes.len() as u64;
883 header.push_str(&format!(
884 "\"{name}\":{{\"dtype\":\"{dtype}\",\"shape\":[{shape_str}],\"data_offsets\":[{offset},{end}]}}"
885 ));
886 offset = end;
887 data.extend_from_slice(bytes);
888 }
889 header.push('}');
890
891 let mut buf = Vec::new();
892 buf.write_u64::<LittleEndian>(header.len() as u64).unwrap();
893 buf.write_all(header.as_bytes()).unwrap();
894 buf.extend_from_slice(&data);
895 buf
896 }
897
898 #[test]
899 fn loads_a_dense_mlp_from_a_real_on_disk_safetensors_shard() {
900 let hidden_dim = 4;
901 let intermediate_dim = 6;
902 let gate = vec![
903 0.1f32, 0.2, -0.3, 0.4, 0.5, -0.6, 0.7, 0.8, -0.9, 1.0, 1.1, -1.2, 0.0, 0.0, 0.0, 0.0,
904 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,
905 ];
906 let gate = &gate[..intermediate_dim * hidden_dim];
907 let up = vec![0.05f32; intermediate_dim * hidden_dim];
908 let down = vec![0.02f32; hidden_dim * intermediate_dim];
909
910 let shard_bytes = build_shard(&[
911 (
912 "model.layers.0.mlp.gate_proj.weight",
913 "BF16",
914 &[intermediate_dim, hidden_dim],
915 bf16_bytes(gate),
916 ),
917 (
918 "model.layers.0.mlp.up_proj.weight",
919 "BF16",
920 &[intermediate_dim, hidden_dim],
921 bf16_bytes(&up),
922 ),
923 (
924 "model.layers.0.mlp.down_proj.weight",
925 "BF16",
926 &[hidden_dim, intermediate_dim],
927 bf16_bytes(&down),
928 ),
929 ]);
930
931 let dir = std::env::temp_dir().join("ferrox_kimi_loader_dense_test");
932 std::fs::create_dir_all(&dir).unwrap();
933 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
934 let index = r#"{"weight_map":{
935 "model.layers.0.mlp.gate_proj.weight":"shard0.safetensors",
936 "model.layers.0.mlp.up_proj.weight":"shard0.safetensors",
937 "model.layers.0.mlp.down_proj.weight":"shard0.safetensors"
938 }}"#;
939 let index_path = dir.join("model.safetensors.index.json");
940 std::fs::write(&index_path, index).unwrap();
941
942 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
943 let weights = load_dense_mlp(&shard, "model.layers.0", hidden_dim, intermediate_dim)
944 .expect("must load dense mlp");
945 std::fs::remove_dir_all(&dir).ok();
946
947 assert_eq!(weights.gate_proj.rows(), intermediate_dim);
948 assert_eq!(weights.gate_proj.cols(), hidden_dim);
949 let x = vec![1.0f32; hidden_dim];
950 let out = weights.forward(&x, 4.0, 25.0);
951 assert_eq!(out.len(), hidden_dim);
952 assert!(out.iter().all(|v| v.is_finite()));
953 }
954
955 #[test]
956 fn a_log_padding_is_truncated_to_num_heads() {
957 let a_log_full: Vec<f32> = (0..8).map(|i| i as f32 * 0.1).collect();
961 let raw: Vec<u8> = a_log_full.iter().flat_map(|v| v.to_le_bytes()).collect();
962
963 let shard_bytes = build_shard(&[("self_attn.A_log", "F32", &[8], raw)]);
964 let dir = std::env::temp_dir().join("ferrox_kimi_loader_alog_test");
965 std::fs::create_dir_all(&dir).unwrap();
966 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
967 let index = r#"{"weight_map":{"self_attn.A_log":"shard0.safetensors"}}"#;
968 let index_path = dir.join("model.safetensors.index.json");
969 std::fs::write(&index_path, index).unwrap();
970
971 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
972 let full = load_f32_vec(&shard, "self_attn.A_log").unwrap();
973 std::fs::remove_dir_all(&dir).ok();
974
975 assert_eq!(full.len(), 8);
976 let truncated = &full[..2];
977 assert_eq!(truncated, &[0.0, 0.1]);
978 }
979
980 fn pseudo_bytes(seed: u32, len: usize) -> Vec<u8> {
985 let mut state = seed.wrapping_mul(2654435761).wrapping_add(1);
986 (0..len)
987 .map(|_| {
988 state = state.wrapping_mul(1103515245).wrapping_add(12345);
989 (state >> 16) as u8
990 })
991 .collect()
992 }
993
994 fn pseudo_scale_bytes(seed: u32, len: usize) -> Vec<u8> {
1005 pseudo_bytes(seed, len)
1006 .into_iter()
1007 .map(|b| b % 180)
1008 .collect()
1009 }
1010
1011 #[test]
1012 fn loads_one_mxfp4_expert_from_a_real_on_disk_safetensors_shard() {
1013 let moe_hidden_dim = 32;
1016 let moe_intermediate_dim = 32;
1017 let expert_prefix = "model.layers.3.block_sparse_moe.experts.0";
1018
1019 let w1_packed = pseudo_bytes(1, moe_intermediate_dim * (moe_hidden_dim / 2));
1020 let w1_scale = pseudo_scale_bytes(2, moe_intermediate_dim * (moe_hidden_dim / 32));
1021 let w2_packed = pseudo_bytes(3, moe_hidden_dim * (moe_intermediate_dim / 2));
1022 let w2_scale = pseudo_scale_bytes(4, moe_hidden_dim * (moe_intermediate_dim / 32));
1023 let w3_packed = pseudo_bytes(5, moe_intermediate_dim * (moe_hidden_dim / 2));
1024 let w3_scale = pseudo_scale_bytes(6, moe_intermediate_dim * (moe_hidden_dim / 32));
1025
1026 let shard_bytes = build_shard(&[
1027 (
1028 &format!("{expert_prefix}.w1.weight_packed"),
1029 "U8",
1030 &[moe_intermediate_dim, moe_hidden_dim / 2],
1031 w1_packed,
1032 ),
1033 (
1034 &format!("{expert_prefix}.w1.weight_scale"),
1035 "U8",
1036 &[moe_intermediate_dim, moe_hidden_dim / 32],
1037 w1_scale,
1038 ),
1039 (
1040 &format!("{expert_prefix}.w2.weight_packed"),
1041 "U8",
1042 &[moe_hidden_dim, moe_intermediate_dim / 2],
1043 w2_packed,
1044 ),
1045 (
1046 &format!("{expert_prefix}.w2.weight_scale"),
1047 "U8",
1048 &[moe_hidden_dim, moe_intermediate_dim / 32],
1049 w2_scale,
1050 ),
1051 (
1052 &format!("{expert_prefix}.w3.weight_packed"),
1053 "U8",
1054 &[moe_intermediate_dim, moe_hidden_dim / 2],
1055 w3_packed,
1056 ),
1057 (
1058 &format!("{expert_prefix}.w3.weight_scale"),
1059 "U8",
1060 &[moe_intermediate_dim, moe_hidden_dim / 32],
1061 w3_scale,
1062 ),
1063 ]);
1064
1065 let dir = std::env::temp_dir().join("ferrox_kimi_loader_mxfp4_test");
1066 std::fs::create_dir_all(&dir).unwrap();
1067 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
1068 let index = format!(
1069 r#"{{"weight_map":{{
1070 "{expert_prefix}.w1.weight_packed":"shard0.safetensors",
1071 "{expert_prefix}.w1.weight_scale":"shard0.safetensors",
1072 "{expert_prefix}.w2.weight_packed":"shard0.safetensors",
1073 "{expert_prefix}.w2.weight_scale":"shard0.safetensors",
1074 "{expert_prefix}.w3.weight_packed":"shard0.safetensors",
1075 "{expert_prefix}.w3.weight_scale":"shard0.safetensors"
1076 }}}}"#
1077 );
1078 let index_path = dir.join("model.safetensors.index.json");
1079 std::fs::write(&index_path, &index).unwrap();
1080
1081 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
1082 let expert = load_kimi_expert(
1083 &shard,
1084 "model.layers.3.block_sparse_moe",
1085 0,
1086 moe_hidden_dim,
1087 moe_intermediate_dim,
1088 )
1089 .expect("must load real MXFP4 expert weights");
1090 std::fs::remove_dir_all(&dir).ok();
1091
1092 assert_eq!(expert.w1.rows(), moe_intermediate_dim);
1093 assert_eq!(expert.w1.cols(), moe_hidden_dim);
1094 assert_eq!(expert.w2.rows(), moe_hidden_dim);
1095 assert_eq!(expert.w2.cols(), moe_intermediate_dim);
1096
1097 let x = vec![0.1f32; moe_hidden_dim];
1098 let out = expert.forward(&x, 4.0, 25.0);
1099 assert_eq!(out.len(), moe_hidden_dim);
1100 assert!(out.iter().all(|v| v.is_finite()));
1101 }
1102
1103 fn build_shard_owned(tensors: Vec<(String, &str, Vec<usize>, Vec<u8>)>) -> Vec<u8> {
1108 let refs: Vec<(&str, &str, &[usize], Vec<u8>)> = tensors
1109 .iter()
1110 .map(|(name, dtype, shape, bytes)| {
1111 (name.as_str(), *dtype, shape.as_slice(), bytes.clone())
1112 })
1113 .collect();
1114 build_shard(&refs)
1115 }
1116
1117 #[test]
1118 fn load_kimi_layer_dispatches_kda_plus_dense_at_a_nonzero_layer_index() {
1119 let hidden_dim = 8;
1120 let kda_num_heads = 2;
1121 let kda_head_dim = 3;
1122 let kda_proj = kda_num_heads * kda_head_dim;
1123 let conv_size = 4;
1124 let dense_intermediate = 5;
1125 let layer_idx = 5;
1126 let prefix = format!("language_model.model.layers.{layer_idx}");
1127
1128 let mut tensors = Vec::new();
1129 let mut push_bf16 = |name: String, shape: Vec<usize>, n: usize| {
1130 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.05f32; n])));
1131 };
1132 push_bf16(
1133 format!("{prefix}.input_layernorm.weight"),
1134 vec![hidden_dim],
1135 hidden_dim,
1136 );
1137 push_bf16(
1138 format!("{prefix}.post_attention_layernorm.weight"),
1139 vec![hidden_dim],
1140 hidden_dim,
1141 );
1142 push_bf16(
1143 format!("{prefix}.self_attention_res_norm.weight"),
1144 vec![hidden_dim],
1145 hidden_dim,
1146 );
1147 push_bf16(
1148 format!("{prefix}.self_attention_res_proj.weight"),
1149 vec![1, hidden_dim],
1150 hidden_dim,
1151 );
1152 push_bf16(
1153 format!("{prefix}.mlp_res_norm.weight"),
1154 vec![hidden_dim],
1155 hidden_dim,
1156 );
1157 push_bf16(
1158 format!("{prefix}.mlp_res_proj.weight"),
1159 vec![1, hidden_dim],
1160 hidden_dim,
1161 );
1162 push_bf16(
1163 format!("{prefix}.self_attn.q_proj.weight"),
1164 vec![kda_proj, hidden_dim],
1165 kda_proj * hidden_dim,
1166 );
1167 push_bf16(
1168 format!("{prefix}.self_attn.k_proj.weight"),
1169 vec![kda_proj, hidden_dim],
1170 kda_proj * hidden_dim,
1171 );
1172 push_bf16(
1173 format!("{prefix}.self_attn.v_proj.weight"),
1174 vec![kda_proj, hidden_dim],
1175 kda_proj * hidden_dim,
1176 );
1177 push_bf16(
1178 format!("{prefix}.self_attn.f_a_proj.weight"),
1179 vec![kda_head_dim, hidden_dim],
1180 kda_head_dim * hidden_dim,
1181 );
1182 push_bf16(
1183 format!("{prefix}.self_attn.f_b_proj.weight"),
1184 vec![kda_proj, kda_head_dim],
1185 kda_proj * kda_head_dim,
1186 );
1187 push_bf16(
1188 format!("{prefix}.self_attn.b_proj.weight"),
1189 vec![kda_num_heads, hidden_dim],
1190 kda_num_heads * hidden_dim,
1191 );
1192 push_bf16(
1193 format!("{prefix}.self_attn.g_proj.weight"),
1194 vec![kda_proj, hidden_dim],
1195 kda_proj * hidden_dim,
1196 );
1197 push_bf16(
1198 format!("{prefix}.self_attn.o_proj.weight"),
1199 vec![hidden_dim, kda_proj],
1200 hidden_dim * kda_proj,
1201 );
1202 push_bf16(
1203 format!("{prefix}.mlp.gate_proj.weight"),
1204 vec![dense_intermediate, hidden_dim],
1205 dense_intermediate * hidden_dim,
1206 );
1207 push_bf16(
1208 format!("{prefix}.mlp.up_proj.weight"),
1209 vec![dense_intermediate, hidden_dim],
1210 dense_intermediate * hidden_dim,
1211 );
1212 push_bf16(
1213 format!("{prefix}.mlp.down_proj.weight"),
1214 vec![hidden_dim, dense_intermediate],
1215 hidden_dim * dense_intermediate,
1216 );
1217
1218 let f32_vec = |v: Vec<f32>| -> Vec<u8> { v.iter().flat_map(|x| x.to_le_bytes()).collect() };
1219 tensors.push((
1220 format!("{prefix}.self_attn.A_log"),
1221 "F32",
1222 vec![kda_num_heads],
1223 f32_vec(vec![0.5; kda_num_heads]),
1224 ));
1225 tensors.push((
1226 format!("{prefix}.self_attn.dt_bias"),
1227 "F32",
1228 vec![kda_proj],
1229 f32_vec(vec![0.1; kda_proj]),
1230 ));
1231 tensors.push((
1232 format!("{prefix}.self_attn.o_norm.weight"),
1233 "F32",
1234 vec![kda_head_dim],
1235 f32_vec(vec![1.0; kda_head_dim]),
1236 ));
1237 for conv_name in ["q_conv1d", "k_conv1d", "v_conv1d"] {
1238 tensors.push((
1239 format!("{prefix}.self_attn.{conv_name}.weight"),
1240 "F32",
1241 vec![kda_proj, 1, conv_size],
1242 f32_vec(vec![0.1; kda_proj * conv_size]),
1243 ));
1244 }
1245
1246 let shard_bytes = build_shard_owned(tensors.clone());
1247 let dir = std::env::temp_dir().join("ferrox_kimi_loader_layer_kda_dense_test");
1248 std::fs::create_dir_all(&dir).unwrap();
1249 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
1250 let map_entries: Vec<String> = tensors
1251 .iter()
1252 .map(|(name, ..)| format!("\"{name}\":\"shard0.safetensors\""))
1253 .collect();
1254 let index = format!("{{\"weight_map\":{{{}}}}}", map_entries.join(","));
1255 let index_path = dir.join("model.safetensors.index.json");
1256 std::fs::write(&index_path, &index).unwrap();
1257
1258 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
1259 let mut hp = KimiRealHparams::real();
1260 hp.hidden_dim = hidden_dim;
1261 hp.kda_num_heads = kda_num_heads;
1262 hp.kda_head_dim = kda_head_dim;
1263 hp.dense_intermediate_dim = dense_intermediate;
1264
1265 let layer = load_kimi_layer(&shard, &hp, LayerAttentionKind::KimiKda, true, layer_idx)
1266 .expect("must load a real KDA+dense layer at a nonzero layer index");
1267 std::fs::remove_dir_all(&dir).ok();
1268
1269 assert!(matches!(
1270 layer.attn,
1271 crate::kimi_decoder::KimiLayerAttention::Kda(_)
1272 ));
1273 assert!(matches!(
1274 layer.ffn,
1275 crate::kimi_decoder::KimiLayerFfn::Dense(_)
1276 ));
1277 assert_eq!(layer.input_layernorm_weight.len(), hidden_dim);
1278 }
1279
1280 #[test]
1281 fn load_kimi_layer_dispatches_mla_plus_latent_moe() {
1282 let hidden_dim = 8;
1283 let num_heads = 1;
1284 let q_lora_rank = 4;
1285 let kv_lora_rank = 4;
1286 let qk_nope_head_dim = 2;
1287 let qk_rope_head_dim = 2;
1288 let v_head_dim = 2;
1289 let q_head_dim = qk_nope_head_dim + qk_rope_head_dim;
1290 let moe_hidden_dim = 32;
1291 let moe_intermediate_dim = 32;
1292 let n_experts = 2;
1293 let num_shared_experts = 1;
1294 let shared_intermediate_dim = moe_intermediate_dim * num_shared_experts;
1295 let layer_idx = 7;
1296 let prefix = format!("language_model.model.layers.{layer_idx}");
1297
1298 let mut tensors: Vec<(String, &str, Vec<usize>, Vec<u8>)> = Vec::new();
1299 let push_bf16 = |tensors: &mut Vec<(String, &str, Vec<usize>, Vec<u8>)>,
1300 name: String,
1301 shape: Vec<usize>,
1302 n: usize| {
1303 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.05f32; n])));
1304 };
1305 push_bf16(
1306 &mut tensors,
1307 format!("{prefix}.input_layernorm.weight"),
1308 vec![hidden_dim],
1309 hidden_dim,
1310 );
1311 push_bf16(
1312 &mut tensors,
1313 format!("{prefix}.post_attention_layernorm.weight"),
1314 vec![hidden_dim],
1315 hidden_dim,
1316 );
1317 push_bf16(
1318 &mut tensors,
1319 format!("{prefix}.self_attention_res_norm.weight"),
1320 vec![hidden_dim],
1321 hidden_dim,
1322 );
1323 push_bf16(
1324 &mut tensors,
1325 format!("{prefix}.self_attention_res_proj.weight"),
1326 vec![1, hidden_dim],
1327 hidden_dim,
1328 );
1329 push_bf16(
1330 &mut tensors,
1331 format!("{prefix}.mlp_res_norm.weight"),
1332 vec![hidden_dim],
1333 hidden_dim,
1334 );
1335 push_bf16(
1336 &mut tensors,
1337 format!("{prefix}.mlp_res_proj.weight"),
1338 vec![1, hidden_dim],
1339 hidden_dim,
1340 );
1341
1342 push_bf16(
1344 &mut tensors,
1345 format!("{prefix}.self_attn.q_a_proj.weight"),
1346 vec![q_lora_rank, hidden_dim],
1347 q_lora_rank * hidden_dim,
1348 );
1349 push_bf16(
1350 &mut tensors,
1351 format!("{prefix}.self_attn.q_a_layernorm.weight"),
1352 vec![q_lora_rank],
1353 q_lora_rank,
1354 );
1355 push_bf16(
1356 &mut tensors,
1357 format!("{prefix}.self_attn.q_b_proj.weight"),
1358 vec![num_heads * q_head_dim, q_lora_rank],
1359 num_heads * q_head_dim * q_lora_rank,
1360 );
1361 push_bf16(
1362 &mut tensors,
1363 format!("{prefix}.self_attn.kv_a_proj_with_mqa.weight"),
1364 vec![kv_lora_rank + qk_rope_head_dim, hidden_dim],
1365 (kv_lora_rank + qk_rope_head_dim) * hidden_dim,
1366 );
1367 push_bf16(
1368 &mut tensors,
1369 format!("{prefix}.self_attn.kv_a_layernorm.weight"),
1370 vec![kv_lora_rank],
1371 kv_lora_rank,
1372 );
1373 push_bf16(
1374 &mut tensors,
1375 format!("{prefix}.self_attn.kv_b_proj.weight"),
1376 vec![num_heads * (qk_nope_head_dim + v_head_dim), kv_lora_rank],
1377 num_heads * (qk_nope_head_dim + v_head_dim) * kv_lora_rank,
1378 );
1379 push_bf16(
1380 &mut tensors,
1381 format!("{prefix}.self_attn.o_proj.weight"),
1382 vec![hidden_dim, num_heads * v_head_dim],
1383 hidden_dim * num_heads * v_head_dim,
1384 );
1385 push_bf16(
1386 &mut tensors,
1387 format!("{prefix}.self_attn.g_proj.weight"),
1388 vec![num_heads * v_head_dim, hidden_dim],
1389 num_heads * v_head_dim * hidden_dim,
1390 );
1391
1392 push_bf16(
1394 &mut tensors,
1395 format!("{prefix}.block_sparse_moe.gate.weight"),
1396 vec![n_experts, hidden_dim],
1397 n_experts * hidden_dim,
1398 );
1399 let bias_bytes: Vec<u8> = vec![0.0f32; n_experts]
1400 .iter()
1401 .flat_map(|v| v.to_le_bytes())
1402 .collect();
1403 tensors.push((
1404 format!("{prefix}.block_sparse_moe.gate.e_score_correction_bias"),
1405 "F32",
1406 vec![n_experts],
1407 bias_bytes,
1408 ));
1409 push_bf16(
1410 &mut tensors,
1411 format!("{prefix}.block_sparse_moe.routed_expert_down_proj.weight"),
1412 vec![moe_hidden_dim, hidden_dim],
1413 moe_hidden_dim * hidden_dim,
1414 );
1415 push_bf16(
1416 &mut tensors,
1417 format!("{prefix}.block_sparse_moe.routed_expert_up_proj.weight"),
1418 vec![hidden_dim, moe_hidden_dim],
1419 hidden_dim * moe_hidden_dim,
1420 );
1421 push_bf16(
1422 &mut tensors,
1423 format!("{prefix}.block_sparse_moe.routed_expert_norm.weight"),
1424 vec![moe_hidden_dim],
1425 moe_hidden_dim,
1426 );
1427 push_bf16(
1428 &mut tensors,
1429 format!("{prefix}.block_sparse_moe.shared_experts.gate_proj.weight"),
1430 vec![shared_intermediate_dim, hidden_dim],
1431 shared_intermediate_dim * hidden_dim,
1432 );
1433 push_bf16(
1434 &mut tensors,
1435 format!("{prefix}.block_sparse_moe.shared_experts.down_proj.weight"),
1436 vec![hidden_dim, shared_intermediate_dim],
1437 hidden_dim * shared_intermediate_dim,
1438 );
1439 push_bf16(
1440 &mut tensors,
1441 format!("{prefix}.block_sparse_moe.shared_experts.up_proj.weight"),
1442 vec![shared_intermediate_dim, hidden_dim],
1443 shared_intermediate_dim * hidden_dim,
1444 );
1445
1446 for e in 0..n_experts {
1447 let expert_prefix = format!("{prefix}.block_sparse_moe.experts.{e}");
1448 let seed_base = (e as u32 + 1) * 10;
1449 tensors.push((
1450 format!("{expert_prefix}.w1.weight_packed"),
1451 "U8",
1452 vec![moe_intermediate_dim, moe_hidden_dim / 2],
1453 pseudo_bytes(seed_base + 1, moe_intermediate_dim * (moe_hidden_dim / 2)),
1454 ));
1455 tensors.push((
1456 format!("{expert_prefix}.w1.weight_scale"),
1457 "U8",
1458 vec![moe_intermediate_dim, moe_hidden_dim / 32],
1459 pseudo_scale_bytes(seed_base + 2, moe_intermediate_dim * (moe_hidden_dim / 32)),
1460 ));
1461 tensors.push((
1462 format!("{expert_prefix}.w2.weight_packed"),
1463 "U8",
1464 vec![moe_hidden_dim, moe_intermediate_dim / 2],
1465 pseudo_bytes(seed_base + 3, moe_hidden_dim * (moe_intermediate_dim / 2)),
1466 ));
1467 tensors.push((
1468 format!("{expert_prefix}.w2.weight_scale"),
1469 "U8",
1470 vec![moe_hidden_dim, moe_intermediate_dim / 32],
1471 pseudo_scale_bytes(seed_base + 4, moe_hidden_dim * (moe_intermediate_dim / 32)),
1472 ));
1473 tensors.push((
1474 format!("{expert_prefix}.w3.weight_packed"),
1475 "U8",
1476 vec![moe_intermediate_dim, moe_hidden_dim / 2],
1477 pseudo_bytes(seed_base + 5, moe_intermediate_dim * (moe_hidden_dim / 2)),
1478 ));
1479 tensors.push((
1480 format!("{expert_prefix}.w3.weight_scale"),
1481 "U8",
1482 vec![moe_intermediate_dim, moe_hidden_dim / 32],
1483 pseudo_scale_bytes(seed_base + 6, moe_intermediate_dim * (moe_hidden_dim / 32)),
1484 ));
1485 }
1486
1487 let shard_bytes = build_shard_owned(tensors.clone());
1488 let dir = std::env::temp_dir().join("ferrox_kimi_loader_layer_mla_moe_test");
1489 std::fs::create_dir_all(&dir).unwrap();
1490 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
1491 let map_entries: Vec<String> = tensors
1492 .iter()
1493 .map(|(name, ..)| format!("\"{name}\":\"shard0.safetensors\""))
1494 .collect();
1495 let index = format!("{{\"weight_map\":{{{}}}}}", map_entries.join(","));
1496 let index_path = dir.join("model.safetensors.index.json");
1497 std::fs::write(&index_path, &index).unwrap();
1498
1499 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
1500 let mut hp = KimiRealHparams::real();
1501 hp.hidden_dim = hidden_dim;
1502 hp.mla_num_heads = num_heads;
1503 hp.mla_q_lora_rank = q_lora_rank;
1504 hp.mla_kv_lora_rank = kv_lora_rank;
1505 hp.mla_qk_nope_head_dim = qk_nope_head_dim;
1506 hp.mla_qk_rope_head_dim = qk_rope_head_dim;
1507 hp.mla_v_head_dim = v_head_dim;
1508 hp.moe_hidden_dim = moe_hidden_dim;
1509 hp.moe_intermediate_dim = moe_intermediate_dim;
1510 hp.n_experts = n_experts;
1511 hp.num_shared_experts = num_shared_experts;
1512
1513 let layer = load_kimi_layer(&shard, &hp, LayerAttentionKind::KimiMla, false, layer_idx)
1514 .expect("must load a real MLA+latent-MoE layer");
1515 std::fs::remove_dir_all(&dir).ok();
1516
1517 assert!(matches!(
1518 layer.attn,
1519 crate::kimi_decoder::KimiLayerAttention::Mla(_)
1520 ));
1521 match &layer.ffn {
1522 crate::kimi_decoder::KimiLayerFfn::Moe(moe) => {
1523 assert_eq!(moe.experts.n_experts(), n_experts);
1524 }
1525 crate::kimi_decoder::KimiLayerFfn::Dense(_) => panic!("expected Moe ffn"),
1526 }
1527 assert_eq!(layer.input_layernorm_weight.len(), hidden_dim);
1528 }
1529
1530 struct SyntheticDims {
1534 hidden_dim: usize,
1535 kda_num_heads: usize,
1536 kda_head_dim: usize,
1537 mla_num_heads: usize,
1538 mla_q_lora_rank: usize,
1539 mla_kv_lora_rank: usize,
1540 mla_qk_nope_head_dim: usize,
1541 mla_qk_rope_head_dim: usize,
1542 mla_v_head_dim: usize,
1543 dense_intermediate_dim: usize,
1544 moe_hidden_dim: usize,
1545 moe_intermediate_dim: usize,
1546 n_experts: usize,
1547 num_shared_experts: usize,
1548 }
1549
1550 #[allow(clippy::too_many_arguments)]
1557 fn push_layer_tensors(
1558 tensors: &mut Vec<(String, &'static str, Vec<usize>, Vec<u8>)>,
1559 layer_idx: usize,
1560 kind: LayerAttentionKind,
1561 is_dense: bool,
1562 d: &SyntheticDims,
1563 ) {
1564 let prefix = format!("language_model.model.layers.{layer_idx}");
1565 let push_bf16 = |tensors: &mut Vec<(String, &'static str, Vec<usize>, Vec<u8>)>,
1566 name: String,
1567 shape: Vec<usize>,
1568 n: usize| {
1569 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.05f32; n])));
1570 };
1571
1572 push_bf16(
1573 tensors,
1574 format!("{prefix}.input_layernorm.weight"),
1575 vec![d.hidden_dim],
1576 d.hidden_dim,
1577 );
1578 push_bf16(
1579 tensors,
1580 format!("{prefix}.post_attention_layernorm.weight"),
1581 vec![d.hidden_dim],
1582 d.hidden_dim,
1583 );
1584 push_bf16(
1585 tensors,
1586 format!("{prefix}.self_attention_res_norm.weight"),
1587 vec![d.hidden_dim],
1588 d.hidden_dim,
1589 );
1590 push_bf16(
1591 tensors,
1592 format!("{prefix}.self_attention_res_proj.weight"),
1593 vec![1, d.hidden_dim],
1594 d.hidden_dim,
1595 );
1596 push_bf16(
1597 tensors,
1598 format!("{prefix}.mlp_res_norm.weight"),
1599 vec![d.hidden_dim],
1600 d.hidden_dim,
1601 );
1602 push_bf16(
1603 tensors,
1604 format!("{prefix}.mlp_res_proj.weight"),
1605 vec![1, d.hidden_dim],
1606 d.hidden_dim,
1607 );
1608
1609 match kind {
1610 LayerAttentionKind::KimiKda => {
1611 let proj = d.kda_num_heads * d.kda_head_dim;
1612 for name in ["q_proj", "k_proj", "v_proj", "g_proj"] {
1613 push_bf16(
1614 tensors,
1615 format!("{prefix}.self_attn.{name}.weight"),
1616 vec![proj, d.hidden_dim],
1617 proj * d.hidden_dim,
1618 );
1619 }
1620 push_bf16(
1621 tensors,
1622 format!("{prefix}.self_attn.f_a_proj.weight"),
1623 vec![d.kda_head_dim, d.hidden_dim],
1624 d.kda_head_dim * d.hidden_dim,
1625 );
1626 push_bf16(
1627 tensors,
1628 format!("{prefix}.self_attn.f_b_proj.weight"),
1629 vec![proj, d.kda_head_dim],
1630 proj * d.kda_head_dim,
1631 );
1632 push_bf16(
1633 tensors,
1634 format!("{prefix}.self_attn.b_proj.weight"),
1635 vec![d.kda_num_heads, d.hidden_dim],
1636 d.kda_num_heads * d.hidden_dim,
1637 );
1638 push_bf16(
1639 tensors,
1640 format!("{prefix}.self_attn.o_proj.weight"),
1641 vec![d.hidden_dim, proj],
1642 d.hidden_dim * proj,
1643 );
1644 let f32_vec =
1645 |v: Vec<f32>| -> Vec<u8> { v.iter().flat_map(|x| x.to_le_bytes()).collect() };
1646 tensors.push((
1647 format!("{prefix}.self_attn.A_log"),
1648 "F32",
1649 vec![d.kda_num_heads],
1650 f32_vec(vec![0.5; d.kda_num_heads]),
1651 ));
1652 tensors.push((
1653 format!("{prefix}.self_attn.dt_bias"),
1654 "F32",
1655 vec![proj],
1656 f32_vec(vec![0.1; proj]),
1657 ));
1658 tensors.push((
1659 format!("{prefix}.self_attn.o_norm.weight"),
1660 "F32",
1661 vec![d.kda_head_dim],
1662 f32_vec(vec![1.0; d.kda_head_dim]),
1663 ));
1664 for conv_name in ["q_conv1d", "k_conv1d", "v_conv1d"] {
1665 tensors.push((
1666 format!("{prefix}.self_attn.{conv_name}.weight"),
1667 "F32",
1668 vec![proj, 1, 4],
1669 f32_vec(vec![0.1; proj * 4]),
1670 ));
1671 }
1672 }
1673 LayerAttentionKind::KimiMla => {
1674 let q_head_dim = d.mla_qk_nope_head_dim + d.mla_qk_rope_head_dim;
1675 push_bf16(
1676 tensors,
1677 format!("{prefix}.self_attn.q_a_proj.weight"),
1678 vec![d.mla_q_lora_rank, d.hidden_dim],
1679 d.mla_q_lora_rank * d.hidden_dim,
1680 );
1681 push_bf16(
1682 tensors,
1683 format!("{prefix}.self_attn.q_a_layernorm.weight"),
1684 vec![d.mla_q_lora_rank],
1685 d.mla_q_lora_rank,
1686 );
1687 push_bf16(
1688 tensors,
1689 format!("{prefix}.self_attn.q_b_proj.weight"),
1690 vec![d.mla_num_heads * q_head_dim, d.mla_q_lora_rank],
1691 d.mla_num_heads * q_head_dim * d.mla_q_lora_rank,
1692 );
1693 push_bf16(
1694 tensors,
1695 format!("{prefix}.self_attn.kv_a_proj_with_mqa.weight"),
1696 vec![d.mla_kv_lora_rank + d.mla_qk_rope_head_dim, d.hidden_dim],
1697 (d.mla_kv_lora_rank + d.mla_qk_rope_head_dim) * d.hidden_dim,
1698 );
1699 push_bf16(
1700 tensors,
1701 format!("{prefix}.self_attn.kv_a_layernorm.weight"),
1702 vec![d.mla_kv_lora_rank],
1703 d.mla_kv_lora_rank,
1704 );
1705 push_bf16(
1706 tensors,
1707 format!("{prefix}.self_attn.kv_b_proj.weight"),
1708 vec![
1709 d.mla_num_heads * (d.mla_qk_nope_head_dim + d.mla_v_head_dim),
1710 d.mla_kv_lora_rank,
1711 ],
1712 d.mla_num_heads
1713 * (d.mla_qk_nope_head_dim + d.mla_v_head_dim)
1714 * d.mla_kv_lora_rank,
1715 );
1716 push_bf16(
1717 tensors,
1718 format!("{prefix}.self_attn.o_proj.weight"),
1719 vec![d.hidden_dim, d.mla_num_heads * d.mla_v_head_dim],
1720 d.hidden_dim * d.mla_num_heads * d.mla_v_head_dim,
1721 );
1722 push_bf16(
1723 tensors,
1724 format!("{prefix}.self_attn.g_proj.weight"),
1725 vec![d.mla_num_heads * d.mla_v_head_dim, d.hidden_dim],
1726 d.mla_num_heads * d.mla_v_head_dim * d.hidden_dim,
1727 );
1728 }
1729 LayerAttentionKind::Gqa => panic!("synthetic checkpoint test never uses Gqa"),
1730 }
1731
1732 if is_dense {
1733 push_bf16(
1734 tensors,
1735 format!("{prefix}.mlp.gate_proj.weight"),
1736 vec![d.dense_intermediate_dim, d.hidden_dim],
1737 d.dense_intermediate_dim * d.hidden_dim,
1738 );
1739 push_bf16(
1740 tensors,
1741 format!("{prefix}.mlp.up_proj.weight"),
1742 vec![d.dense_intermediate_dim, d.hidden_dim],
1743 d.dense_intermediate_dim * d.hidden_dim,
1744 );
1745 push_bf16(
1746 tensors,
1747 format!("{prefix}.mlp.down_proj.weight"),
1748 vec![d.hidden_dim, d.dense_intermediate_dim],
1749 d.hidden_dim * d.dense_intermediate_dim,
1750 );
1751 } else {
1752 let shared_intermediate_dim = d.moe_intermediate_dim * d.num_shared_experts;
1753 push_bf16(
1754 tensors,
1755 format!("{prefix}.block_sparse_moe.gate.weight"),
1756 vec![d.n_experts, d.hidden_dim],
1757 d.n_experts * d.hidden_dim,
1758 );
1759 let bias_bytes: Vec<u8> = vec![0.0f32; d.n_experts]
1760 .iter()
1761 .flat_map(|v| v.to_le_bytes())
1762 .collect();
1763 tensors.push((
1764 format!("{prefix}.block_sparse_moe.gate.e_score_correction_bias"),
1765 "F32",
1766 vec![d.n_experts],
1767 bias_bytes,
1768 ));
1769 push_bf16(
1770 tensors,
1771 format!("{prefix}.block_sparse_moe.routed_expert_down_proj.weight"),
1772 vec![d.moe_hidden_dim, d.hidden_dim],
1773 d.moe_hidden_dim * d.hidden_dim,
1774 );
1775 push_bf16(
1776 tensors,
1777 format!("{prefix}.block_sparse_moe.routed_expert_up_proj.weight"),
1778 vec![d.hidden_dim, d.moe_hidden_dim],
1779 d.hidden_dim * d.moe_hidden_dim,
1780 );
1781 push_bf16(
1782 tensors,
1783 format!("{prefix}.block_sparse_moe.routed_expert_norm.weight"),
1784 vec![d.moe_hidden_dim],
1785 d.moe_hidden_dim,
1786 );
1787 push_bf16(
1788 tensors,
1789 format!("{prefix}.block_sparse_moe.shared_experts.gate_proj.weight"),
1790 vec![shared_intermediate_dim, d.hidden_dim],
1791 shared_intermediate_dim * d.hidden_dim,
1792 );
1793 push_bf16(
1794 tensors,
1795 format!("{prefix}.block_sparse_moe.shared_experts.down_proj.weight"),
1796 vec![d.hidden_dim, shared_intermediate_dim],
1797 d.hidden_dim * shared_intermediate_dim,
1798 );
1799 push_bf16(
1800 tensors,
1801 format!("{prefix}.block_sparse_moe.shared_experts.up_proj.weight"),
1802 vec![shared_intermediate_dim, d.hidden_dim],
1803 shared_intermediate_dim * d.hidden_dim,
1804 );
1805
1806 for e in 0..d.n_experts {
1807 let expert_prefix = format!("{prefix}.block_sparse_moe.experts.{e}");
1808 let seed_base = (layer_idx as u32 * 100) + (e as u32 + 1) * 10;
1809 tensors.push((
1810 format!("{expert_prefix}.w1.weight_packed"),
1811 "U8",
1812 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 2],
1813 pseudo_bytes(
1814 seed_base + 1,
1815 d.moe_intermediate_dim * (d.moe_hidden_dim / 2),
1816 ),
1817 ));
1818 tensors.push((
1819 format!("{expert_prefix}.w1.weight_scale"),
1820 "U8",
1821 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 32],
1822 pseudo_scale_bytes(
1823 seed_base + 2,
1824 d.moe_intermediate_dim * (d.moe_hidden_dim / 32),
1825 ),
1826 ));
1827 tensors.push((
1828 format!("{expert_prefix}.w2.weight_packed"),
1829 "U8",
1830 vec![d.moe_hidden_dim, d.moe_intermediate_dim / 2],
1831 pseudo_bytes(
1832 seed_base + 3,
1833 d.moe_hidden_dim * (d.moe_intermediate_dim / 2),
1834 ),
1835 ));
1836 tensors.push((
1837 format!("{expert_prefix}.w2.weight_scale"),
1838 "U8",
1839 vec![d.moe_hidden_dim, d.moe_intermediate_dim / 32],
1840 pseudo_scale_bytes(
1841 seed_base + 4,
1842 d.moe_hidden_dim * (d.moe_intermediate_dim / 32),
1843 ),
1844 ));
1845 tensors.push((
1846 format!("{expert_prefix}.w3.weight_packed"),
1847 "U8",
1848 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 2],
1849 pseudo_bytes(
1850 seed_base + 5,
1851 d.moe_intermediate_dim * (d.moe_hidden_dim / 2),
1852 ),
1853 ));
1854 tensors.push((
1855 format!("{expert_prefix}.w3.weight_scale"),
1856 "U8",
1857 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 32],
1858 pseudo_scale_bytes(
1859 seed_base + 6,
1860 d.moe_intermediate_dim * (d.moe_hidden_dim / 32),
1861 ),
1862 ));
1863 }
1864 }
1865 }
1866
1867 fn build_synthetic_full_checkpoint(
1871 dir_name: &str,
1872 ) -> (
1873 std::path::PathBuf,
1874 ShardedSafetensors,
1875 crate::config::ModelConfig,
1876 KimiRealHparams,
1877 ) {
1878 let d = SyntheticDims {
1879 hidden_dim: 8,
1880 kda_num_heads: 2,
1881 kda_head_dim: 3,
1882 mla_num_heads: 1,
1883 mla_q_lora_rank: 4,
1884 mla_kv_lora_rank: 4,
1885 mla_qk_nope_head_dim: 2,
1886 mla_qk_rope_head_dim: 2,
1887 mla_v_head_dim: 2,
1888 dense_intermediate_dim: 5,
1889 moe_hidden_dim: 32,
1890 moe_intermediate_dim: 32,
1891 n_experts: 2,
1892 num_shared_experts: 1,
1893 };
1894 let vocab_size = 6;
1895
1896 let model_cfg = crate::config::ModelConfig {
1901 name: "synthetic-kimi-test",
1902 n_layers: 3,
1903 hidden_dim: d.hidden_dim,
1904 n_heads: 1,
1905 n_kv_heads: 1,
1906 head_dim: 4,
1907 vocab_size,
1908 rope_theta: 10000.0,
1909 rms_norm_eps: 1e-5,
1910 sliding_window: None,
1911 moe: ferrox_moe::MoeLayerConfig {
1912 n_experts: d.n_experts,
1913 n_experts_active: d.n_experts,
1914 n_shared_experts: d.num_shared_experts,
1915 hidden_dim: d.hidden_dim,
1916 expert_ffn_dim: d.moe_intermediate_dim,
1917 gating: ferrox_moe::GatingFunction::Sigmoid,
1918 norm_topk_prob: true,
1919 expert_group_count: None,
1920 expert_group_used_count: None,
1921 },
1922 n_dense_leading_layers: 1,
1923 attention: crate::config::AttentionKind::KimiHybrid(
1924 crate::config::KimiHybridAttention {
1925 kda_layers: vec![1, 2],
1926 full_attn_layers: vec![3],
1927 mla: crate::config::MlaConfig {
1928 num_heads: d.mla_num_heads,
1929 q_lora_rank: d.mla_q_lora_rank,
1930 kv_lora_rank: d.mla_kv_lora_rank,
1931 qk_nope_head_dim: d.mla_qk_nope_head_dim,
1932 qk_rope_head_dim: d.mla_qk_rope_head_dim,
1933 v_head_dim: d.mla_v_head_dim,
1934 use_output_gate: true,
1935 rope: None,
1936 },
1937 kda: crate::config::KdaConfig {
1938 num_heads: d.kda_num_heads,
1939 head_dim: d.kda_head_dim,
1940 short_conv_kernel_size: 4,
1941 gate_lower_bound: -5.0,
1942 use_full_rank_gate: true,
1943 },
1944 },
1945 ),
1946 rope_freqs: None,
1947 rope_attn_factor: 1.0,
1948 rope_dim: None,
1949 rope_freqs_long: None,
1950 rope_freqs_short: None,
1951 rope_orig_ctx: None,
1952 rope_layout: crate::config::RopeLayout::Neox,
1953 qk_norm_style: crate::capability::QkNormStyle::WholeVector,
1954 swa_pattern: None,
1955 attn_logit_softcap: None,
1956 final_logit_softcap: None,
1957 embedding_scale: None,
1958 attention_scale: None,
1959 rope_theta_swa: None,
1960 ffn_activation: crate::config::FfnActivation::Swiglu,
1961 best_effort_fields: &["synthetic test config, not a real preset"],
1962 };
1963
1964 let mut tensors: Vec<(String, &'static str, Vec<usize>, Vec<u8>)> = Vec::new();
1965 push_layer_tensors(&mut tensors, 0, LayerAttentionKind::KimiKda, true, &d);
1966 push_layer_tensors(&mut tensors, 1, LayerAttentionKind::KimiKda, false, &d);
1967 push_layer_tensors(&mut tensors, 2, LayerAttentionKind::KimiMla, false, &d);
1968
1969 let mut push_bf16_top = |name: String, shape: Vec<usize>, n: usize| {
1970 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.02f32; n])));
1971 };
1972 push_bf16_top(
1973 "language_model.model.embed_tokens.weight".to_string(),
1974 vec![vocab_size, d.hidden_dim],
1975 vocab_size * d.hidden_dim,
1976 );
1977 push_bf16_top(
1978 "language_model.lm_head.weight".to_string(),
1979 vec![vocab_size, d.hidden_dim],
1980 vocab_size * d.hidden_dim,
1981 );
1982 push_bf16_top(
1983 "language_model.model.norm.weight".to_string(),
1984 vec![d.hidden_dim],
1985 d.hidden_dim,
1986 );
1987 push_bf16_top(
1988 "language_model.model.output_attn_res_norm.weight".to_string(),
1989 vec![d.hidden_dim],
1990 d.hidden_dim,
1991 );
1992 push_bf16_top(
1993 "language_model.model.output_attn_res_proj.weight".to_string(),
1994 vec![1, d.hidden_dim],
1995 d.hidden_dim,
1996 );
1997
1998 let shard_bytes = build_shard_owned(tensors.clone());
1999 let dir = std::env::temp_dir().join(dir_name);
2000 std::fs::create_dir_all(&dir).unwrap();
2001 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
2002 let map_entries: Vec<String> = tensors
2003 .iter()
2004 .map(|(name, ..)| format!("\"{name}\":\"shard0.safetensors\""))
2005 .collect();
2006 let index = format!("{{\"weight_map\":{{{}}}}}", map_entries.join(","));
2007 let index_path = dir.join("model.safetensors.index.json");
2008 std::fs::write(&index_path, &index).unwrap();
2009
2010 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
2011 let hp = KimiRealHparams {
2012 hidden_dim: d.hidden_dim,
2013 kda_num_heads: d.kda_num_heads,
2014 kda_head_dim: d.kda_head_dim,
2015 mla_num_heads: d.mla_num_heads,
2016 mla_q_lora_rank: d.mla_q_lora_rank,
2017 mla_kv_lora_rank: d.mla_kv_lora_rank,
2018 mla_qk_nope_head_dim: d.mla_qk_nope_head_dim,
2019 mla_qk_rope_head_dim: d.mla_qk_rope_head_dim,
2020 mla_v_head_dim: d.mla_v_head_dim,
2021 dense_intermediate_dim: d.dense_intermediate_dim,
2022 moe_hidden_dim: d.moe_hidden_dim,
2023 moe_intermediate_dim: d.moe_intermediate_dim,
2024 n_experts: d.n_experts,
2025 num_shared_experts: d.num_shared_experts,
2026 };
2027
2028 (dir, shard, model_cfg, hp)
2029 }
2030
2031 #[test]
2037 fn store_backed_kimi_experts_produce_bit_identical_outputs() {
2038 let (dir, shard, model_cfg, hp) =
2039 build_synthetic_full_checkpoint("ferrox_kimi_store_equivalence_test");
2040
2041 let eager = load_kimi_checkpoint(&shard, &model_cfg, &hp).expect("eager load");
2042 let mla_cfg = crate::config::MlaConfig {
2043 num_heads: hp.mla_num_heads,
2044 q_lora_rank: hp.mla_q_lora_rank,
2045 kv_lora_rank: hp.mla_kv_lora_rank,
2046 qk_nope_head_dim: hp.mla_qk_nope_head_dim,
2047 qk_rope_head_dim: hp.mla_qk_rope_head_dim,
2048 v_head_dim: hp.mla_v_head_dim,
2049 use_output_gate: true,
2050 rope: None,
2051 };
2052 let kda_cfg = crate::config::KdaConfig {
2053 num_heads: hp.kda_num_heads,
2054 head_dim: hp.kda_head_dim,
2055 short_conv_kernel_size: 4,
2056 gate_lower_bound: -5.0,
2057 use_full_rank_gate: true,
2058 };
2059 let dec_cfg = crate::kimi_decoder::KimiDecoderConfig {
2060 attn_res_block_size: 12,
2061 rms_norm_eps: 1e-5,
2062 situ_beta: 4.0,
2063 situ_linear_beta: 25.0,
2064 moe: crate::latent_moe::KimiMoeConfig {
2065 n_experts_active: hp.n_experts,
2066 moe_renormalize: true,
2067 routed_scaling_factor: 1.0,
2068 situ_beta: 4.0,
2069 situ_linear_beta: 25.0,
2070 rms_norm_eps: 1e-5,
2071 },
2072 };
2073
2074 for budget in [64 * 1024 * 1024u64, 1u64] {
2075 let stored =
2076 load_kimi_checkpoint_with_expert_cache(&shard, &model_cfg, &hp, Some(budget))
2077 .expect("store-backed load");
2078 let mut state_a = crate::kimi_decoder::KimiDecodeState::new(&eager, &kda_cfg);
2079 let mut state_b = crate::kimi_decoder::KimiDecodeState::new(&stored, &kda_cfg);
2080 for &tok in &[1usize, 3, 0, 2] {
2081 let a = crate::kimi_decoder::kimi_forward_token(
2082 &eager,
2083 &dec_cfg,
2084 &mla_cfg,
2085 &kda_cfg,
2086 tok,
2087 &mut state_a,
2088 );
2089 let b = crate::kimi_decoder::kimi_forward_token(
2090 &stored,
2091 &dec_cfg,
2092 &mla_cfg,
2093 &kda_cfg,
2094 tok,
2095 &mut state_b,
2096 );
2097 assert_eq!(
2098 a, b,
2099 "budget={budget}: store-backed Kimi output must be bit-identical"
2100 );
2101 }
2102 }
2103 std::fs::remove_dir_all(&dir).ok();
2104 }
2105
2106 #[test]
2107 fn load_kimi_checkpoint_assembles_every_real_layer_kind() {
2108 let (dir, shard, model_cfg, hp) =
2109 build_synthetic_full_checkpoint("ferrox_kimi_loader_full_checkpoint_test");
2110
2111 let weights = load_kimi_checkpoint(&shard, &model_cfg, &hp)
2112 .expect("must assemble a complete synthetic checkpoint");
2113 std::fs::remove_dir_all(&dir).ok();
2114
2115 assert_eq!(weights.layers.len(), 3);
2116 assert!(matches!(
2117 weights.layers[0].ffn,
2118 crate::kimi_decoder::KimiLayerFfn::Dense(_)
2119 ));
2120 assert!(matches!(
2121 weights.layers[0].attn,
2122 crate::kimi_decoder::KimiLayerAttention::Kda(_)
2123 ));
2124 assert!(matches!(
2125 weights.layers[1].ffn,
2126 crate::kimi_decoder::KimiLayerFfn::Moe(_)
2127 ));
2128 assert!(matches!(
2129 weights.layers[1].attn,
2130 crate::kimi_decoder::KimiLayerAttention::Kda(_)
2131 ));
2132 assert!(matches!(
2133 weights.layers[2].ffn,
2134 crate::kimi_decoder::KimiLayerFfn::Moe(_)
2135 ));
2136 assert!(matches!(
2137 weights.layers[2].attn,
2138 crate::kimi_decoder::KimiLayerAttention::Mla(_)
2139 ));
2140 assert_eq!(weights.embedding.rows(), model_cfg.vocab_size);
2141 assert_eq!(weights.embedding.cols(), hp.hidden_dim);
2142 assert_eq!(weights.output_head.rows(), model_cfg.vocab_size);
2143 assert_eq!(weights.final_norm_weight.len(), hp.hidden_dim);
2144
2145 let mla_cfg = crate::config::MlaConfig {
2149 num_heads: hp.mla_num_heads,
2150 q_lora_rank: hp.mla_q_lora_rank,
2151 kv_lora_rank: hp.mla_kv_lora_rank,
2152 qk_nope_head_dim: hp.mla_qk_nope_head_dim,
2153 qk_rope_head_dim: hp.mla_qk_rope_head_dim,
2154 v_head_dim: hp.mla_v_head_dim,
2155 use_output_gate: true,
2156 rope: None,
2157 };
2158 let kda_cfg = crate::config::KdaConfig {
2159 num_heads: hp.kda_num_heads,
2160 head_dim: hp.kda_head_dim,
2161 short_conv_kernel_size: 4,
2162 gate_lower_bound: -5.0,
2163 use_full_rank_gate: true,
2164 };
2165 let decoder_cfg = crate::kimi_decoder::KimiDecoderConfig {
2166 attn_res_block_size: 12,
2167 rms_norm_eps: 1e-5,
2168 situ_beta: 4.0,
2169 situ_linear_beta: 25.0,
2170 moe: crate::latent_moe::KimiMoeConfig {
2171 n_experts_active: hp.n_experts,
2172 moe_renormalize: true,
2173 routed_scaling_factor: 1.0,
2174 situ_beta: 4.0,
2175 situ_linear_beta: 25.0,
2176 rms_norm_eps: 1e-5,
2177 },
2178 };
2179 let mut state = crate::kimi_decoder::KimiDecodeState::new(&weights, &kda_cfg);
2180 let logits = crate::kimi_decoder::kimi_forward_token(
2181 &weights,
2182 &decoder_cfg,
2183 &mla_cfg,
2184 &kda_cfg,
2185 0,
2186 &mut state,
2187 );
2188 assert_eq!(logits.len(), model_cfg.vocab_size);
2189 assert!(logits.iter().all(|v| v.is_finite()));
2190 }
2191}