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.as_chunks::<4>().0 {
80 out.push(f32::from_le_bytes(*chunk));
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(format!(
932 "ferrox_kimi_loader_dense_test_{}",
933 std::process::id()
934 ));
935 std::fs::create_dir_all(&dir).unwrap();
936 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
937 let index = r#"{"weight_map":{
938 "model.layers.0.mlp.gate_proj.weight":"shard0.safetensors",
939 "model.layers.0.mlp.up_proj.weight":"shard0.safetensors",
940 "model.layers.0.mlp.down_proj.weight":"shard0.safetensors"
941 }}"#;
942 let index_path = dir.join("model.safetensors.index.json");
943 std::fs::write(&index_path, index).unwrap();
944
945 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
946 let weights = load_dense_mlp(&shard, "model.layers.0", hidden_dim, intermediate_dim)
947 .expect("must load dense mlp");
948 std::fs::remove_dir_all(&dir).ok();
949
950 assert_eq!(weights.gate_proj.rows(), intermediate_dim);
951 assert_eq!(weights.gate_proj.cols(), hidden_dim);
952 let x = vec![1.0f32; hidden_dim];
953 let out = weights.forward(&x, 4.0, 25.0);
954 assert_eq!(out.len(), hidden_dim);
955 assert!(out.iter().all(|v| v.is_finite()));
956 }
957
958 #[test]
959 fn a_log_padding_is_truncated_to_num_heads() {
960 let a_log_full: Vec<f32> = (0..8).map(|i| i as f32 * 0.1).collect();
964 let raw: Vec<u8> = a_log_full.iter().flat_map(|v| v.to_le_bytes()).collect();
965
966 let shard_bytes = build_shard(&[("self_attn.A_log", "F32", &[8], raw)]);
967 let dir = std::env::temp_dir().join(format!(
968 "ferrox_kimi_loader_alog_test_{}",
969 std::process::id()
970 ));
971 std::fs::create_dir_all(&dir).unwrap();
972 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
973 let index = r#"{"weight_map":{"self_attn.A_log":"shard0.safetensors"}}"#;
974 let index_path = dir.join("model.safetensors.index.json");
975 std::fs::write(&index_path, index).unwrap();
976
977 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
978 let full = load_f32_vec(&shard, "self_attn.A_log").unwrap();
979 std::fs::remove_dir_all(&dir).ok();
980
981 assert_eq!(full.len(), 8);
982 let truncated = &full[..2];
983 assert_eq!(truncated, &[0.0, 0.1]);
984 }
985
986 fn pseudo_bytes(seed: u32, len: usize) -> Vec<u8> {
991 let mut state = seed.wrapping_mul(2654435761).wrapping_add(1);
992 (0..len)
993 .map(|_| {
994 state = state.wrapping_mul(1103515245).wrapping_add(12345);
995 (state >> 16) as u8
996 })
997 .collect()
998 }
999
1000 fn pseudo_scale_bytes(seed: u32, len: usize) -> Vec<u8> {
1011 pseudo_bytes(seed, len)
1012 .into_iter()
1013 .map(|b| b % 180)
1014 .collect()
1015 }
1016
1017 #[test]
1018 fn loads_one_mxfp4_expert_from_a_real_on_disk_safetensors_shard() {
1019 let moe_hidden_dim = 32;
1022 let moe_intermediate_dim = 32;
1023 let expert_prefix = "model.layers.3.block_sparse_moe.experts.0";
1024
1025 let w1_packed = pseudo_bytes(1, moe_intermediate_dim * (moe_hidden_dim / 2));
1026 let w1_scale = pseudo_scale_bytes(2, moe_intermediate_dim * (moe_hidden_dim / 32));
1027 let w2_packed = pseudo_bytes(3, moe_hidden_dim * (moe_intermediate_dim / 2));
1028 let w2_scale = pseudo_scale_bytes(4, moe_hidden_dim * (moe_intermediate_dim / 32));
1029 let w3_packed = pseudo_bytes(5, moe_intermediate_dim * (moe_hidden_dim / 2));
1030 let w3_scale = pseudo_scale_bytes(6, moe_intermediate_dim * (moe_hidden_dim / 32));
1031
1032 let shard_bytes = build_shard(&[
1033 (
1034 &format!("{expert_prefix}.w1.weight_packed"),
1035 "U8",
1036 &[moe_intermediate_dim, moe_hidden_dim / 2],
1037 w1_packed,
1038 ),
1039 (
1040 &format!("{expert_prefix}.w1.weight_scale"),
1041 "U8",
1042 &[moe_intermediate_dim, moe_hidden_dim / 32],
1043 w1_scale,
1044 ),
1045 (
1046 &format!("{expert_prefix}.w2.weight_packed"),
1047 "U8",
1048 &[moe_hidden_dim, moe_intermediate_dim / 2],
1049 w2_packed,
1050 ),
1051 (
1052 &format!("{expert_prefix}.w2.weight_scale"),
1053 "U8",
1054 &[moe_hidden_dim, moe_intermediate_dim / 32],
1055 w2_scale,
1056 ),
1057 (
1058 &format!("{expert_prefix}.w3.weight_packed"),
1059 "U8",
1060 &[moe_intermediate_dim, moe_hidden_dim / 2],
1061 w3_packed,
1062 ),
1063 (
1064 &format!("{expert_prefix}.w3.weight_scale"),
1065 "U8",
1066 &[moe_intermediate_dim, moe_hidden_dim / 32],
1067 w3_scale,
1068 ),
1069 ]);
1070
1071 let dir = std::env::temp_dir().join(format!(
1072 "ferrox_kimi_loader_mxfp4_test_{}",
1073 std::process::id()
1074 ));
1075 std::fs::create_dir_all(&dir).unwrap();
1076 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
1077 let index = format!(
1078 r#"{{"weight_map":{{
1079 "{expert_prefix}.w1.weight_packed":"shard0.safetensors",
1080 "{expert_prefix}.w1.weight_scale":"shard0.safetensors",
1081 "{expert_prefix}.w2.weight_packed":"shard0.safetensors",
1082 "{expert_prefix}.w2.weight_scale":"shard0.safetensors",
1083 "{expert_prefix}.w3.weight_packed":"shard0.safetensors",
1084 "{expert_prefix}.w3.weight_scale":"shard0.safetensors"
1085 }}}}"#
1086 );
1087 let index_path = dir.join("model.safetensors.index.json");
1088 std::fs::write(&index_path, &index).unwrap();
1089
1090 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
1091 let expert = load_kimi_expert(
1092 &shard,
1093 "model.layers.3.block_sparse_moe",
1094 0,
1095 moe_hidden_dim,
1096 moe_intermediate_dim,
1097 )
1098 .expect("must load real MXFP4 expert weights");
1099 std::fs::remove_dir_all(&dir).ok();
1100
1101 assert_eq!(expert.w1.rows(), moe_intermediate_dim);
1102 assert_eq!(expert.w1.cols(), moe_hidden_dim);
1103 assert_eq!(expert.w2.rows(), moe_hidden_dim);
1104 assert_eq!(expert.w2.cols(), moe_intermediate_dim);
1105
1106 let x = vec![0.1f32; moe_hidden_dim];
1107 let out = expert.forward(&x, 4.0, 25.0);
1108 assert_eq!(out.len(), moe_hidden_dim);
1109 assert!(out.iter().all(|v| v.is_finite()));
1110 }
1111
1112 fn build_shard_owned(tensors: Vec<(String, &str, Vec<usize>, Vec<u8>)>) -> Vec<u8> {
1117 let refs: Vec<(&str, &str, &[usize], Vec<u8>)> = tensors
1118 .iter()
1119 .map(|(name, dtype, shape, bytes)| {
1120 (name.as_str(), *dtype, shape.as_slice(), bytes.clone())
1121 })
1122 .collect();
1123 build_shard(&refs)
1124 }
1125
1126 #[test]
1127 fn load_kimi_layer_dispatches_kda_plus_dense_at_a_nonzero_layer_index() {
1128 let hidden_dim = 8;
1129 let kda_num_heads = 2;
1130 let kda_head_dim = 3;
1131 let kda_proj = kda_num_heads * kda_head_dim;
1132 let conv_size = 4;
1133 let dense_intermediate = 5;
1134 let layer_idx = 5;
1135 let prefix = format!("language_model.model.layers.{layer_idx}");
1136
1137 let mut tensors = Vec::new();
1138 let mut push_bf16 = |name: String, shape: Vec<usize>, n: usize| {
1139 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.05f32; n])));
1140 };
1141 push_bf16(
1142 format!("{prefix}.input_layernorm.weight"),
1143 vec![hidden_dim],
1144 hidden_dim,
1145 );
1146 push_bf16(
1147 format!("{prefix}.post_attention_layernorm.weight"),
1148 vec![hidden_dim],
1149 hidden_dim,
1150 );
1151 push_bf16(
1152 format!("{prefix}.self_attention_res_norm.weight"),
1153 vec![hidden_dim],
1154 hidden_dim,
1155 );
1156 push_bf16(
1157 format!("{prefix}.self_attention_res_proj.weight"),
1158 vec![1, hidden_dim],
1159 hidden_dim,
1160 );
1161 push_bf16(
1162 format!("{prefix}.mlp_res_norm.weight"),
1163 vec![hidden_dim],
1164 hidden_dim,
1165 );
1166 push_bf16(
1167 format!("{prefix}.mlp_res_proj.weight"),
1168 vec![1, hidden_dim],
1169 hidden_dim,
1170 );
1171 push_bf16(
1172 format!("{prefix}.self_attn.q_proj.weight"),
1173 vec![kda_proj, hidden_dim],
1174 kda_proj * hidden_dim,
1175 );
1176 push_bf16(
1177 format!("{prefix}.self_attn.k_proj.weight"),
1178 vec![kda_proj, hidden_dim],
1179 kda_proj * hidden_dim,
1180 );
1181 push_bf16(
1182 format!("{prefix}.self_attn.v_proj.weight"),
1183 vec![kda_proj, hidden_dim],
1184 kda_proj * hidden_dim,
1185 );
1186 push_bf16(
1187 format!("{prefix}.self_attn.f_a_proj.weight"),
1188 vec![kda_head_dim, hidden_dim],
1189 kda_head_dim * hidden_dim,
1190 );
1191 push_bf16(
1192 format!("{prefix}.self_attn.f_b_proj.weight"),
1193 vec![kda_proj, kda_head_dim],
1194 kda_proj * kda_head_dim,
1195 );
1196 push_bf16(
1197 format!("{prefix}.self_attn.b_proj.weight"),
1198 vec![kda_num_heads, hidden_dim],
1199 kda_num_heads * hidden_dim,
1200 );
1201 push_bf16(
1202 format!("{prefix}.self_attn.g_proj.weight"),
1203 vec![kda_proj, hidden_dim],
1204 kda_proj * hidden_dim,
1205 );
1206 push_bf16(
1207 format!("{prefix}.self_attn.o_proj.weight"),
1208 vec![hidden_dim, kda_proj],
1209 hidden_dim * kda_proj,
1210 );
1211 push_bf16(
1212 format!("{prefix}.mlp.gate_proj.weight"),
1213 vec![dense_intermediate, hidden_dim],
1214 dense_intermediate * hidden_dim,
1215 );
1216 push_bf16(
1217 format!("{prefix}.mlp.up_proj.weight"),
1218 vec![dense_intermediate, hidden_dim],
1219 dense_intermediate * hidden_dim,
1220 );
1221 push_bf16(
1222 format!("{prefix}.mlp.down_proj.weight"),
1223 vec![hidden_dim, dense_intermediate],
1224 hidden_dim * dense_intermediate,
1225 );
1226
1227 let f32_vec = |v: Vec<f32>| -> Vec<u8> { v.iter().flat_map(|x| x.to_le_bytes()).collect() };
1228 tensors.push((
1229 format!("{prefix}.self_attn.A_log"),
1230 "F32",
1231 vec![kda_num_heads],
1232 f32_vec(vec![0.5; kda_num_heads]),
1233 ));
1234 tensors.push((
1235 format!("{prefix}.self_attn.dt_bias"),
1236 "F32",
1237 vec![kda_proj],
1238 f32_vec(vec![0.1; kda_proj]),
1239 ));
1240 tensors.push((
1241 format!("{prefix}.self_attn.o_norm.weight"),
1242 "F32",
1243 vec![kda_head_dim],
1244 f32_vec(vec![1.0; kda_head_dim]),
1245 ));
1246 for conv_name in ["q_conv1d", "k_conv1d", "v_conv1d"] {
1247 tensors.push((
1248 format!("{prefix}.self_attn.{conv_name}.weight"),
1249 "F32",
1250 vec![kda_proj, 1, conv_size],
1251 f32_vec(vec![0.1; kda_proj * conv_size]),
1252 ));
1253 }
1254
1255 let shard_bytes = build_shard_owned(tensors.clone());
1256 let dir = std::env::temp_dir().join(format!(
1257 "ferrox_kimi_loader_layer_kda_dense_test_{}",
1258 std::process::id()
1259 ));
1260 std::fs::create_dir_all(&dir).unwrap();
1261 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
1262 let map_entries: Vec<String> = tensors
1263 .iter()
1264 .map(|(name, ..)| format!("\"{name}\":\"shard0.safetensors\""))
1265 .collect();
1266 let index = format!("{{\"weight_map\":{{{}}}}}", map_entries.join(","));
1267 let index_path = dir.join("model.safetensors.index.json");
1268 std::fs::write(&index_path, &index).unwrap();
1269
1270 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
1271 let mut hp = KimiRealHparams::real();
1272 hp.hidden_dim = hidden_dim;
1273 hp.kda_num_heads = kda_num_heads;
1274 hp.kda_head_dim = kda_head_dim;
1275 hp.dense_intermediate_dim = dense_intermediate;
1276
1277 let layer = load_kimi_layer(&shard, &hp, LayerAttentionKind::KimiKda, true, layer_idx)
1278 .expect("must load a real KDA+dense layer at a nonzero layer index");
1279 std::fs::remove_dir_all(&dir).ok();
1280
1281 assert!(matches!(
1282 layer.attn,
1283 crate::kimi_decoder::KimiLayerAttention::Kda(_)
1284 ));
1285 assert!(matches!(
1286 layer.ffn,
1287 crate::kimi_decoder::KimiLayerFfn::Dense(_)
1288 ));
1289 assert_eq!(layer.input_layernorm_weight.len(), hidden_dim);
1290 }
1291
1292 #[test]
1293 fn load_kimi_layer_dispatches_mla_plus_latent_moe() {
1294 let hidden_dim = 8;
1295 let num_heads = 1;
1296 let q_lora_rank = 4;
1297 let kv_lora_rank = 4;
1298 let qk_nope_head_dim = 2;
1299 let qk_rope_head_dim = 2;
1300 let v_head_dim = 2;
1301 let q_head_dim = qk_nope_head_dim + qk_rope_head_dim;
1302 let moe_hidden_dim = 32;
1303 let moe_intermediate_dim = 32;
1304 let n_experts = 2;
1305 let num_shared_experts = 1;
1306 let shared_intermediate_dim = moe_intermediate_dim * num_shared_experts;
1307 let layer_idx = 7;
1308 let prefix = format!("language_model.model.layers.{layer_idx}");
1309
1310 let mut tensors: Vec<(String, &str, Vec<usize>, Vec<u8>)> = Vec::new();
1311 let push_bf16 = |tensors: &mut Vec<(String, &str, Vec<usize>, Vec<u8>)>,
1312 name: String,
1313 shape: Vec<usize>,
1314 n: usize| {
1315 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.05f32; n])));
1316 };
1317 push_bf16(
1318 &mut tensors,
1319 format!("{prefix}.input_layernorm.weight"),
1320 vec![hidden_dim],
1321 hidden_dim,
1322 );
1323 push_bf16(
1324 &mut tensors,
1325 format!("{prefix}.post_attention_layernorm.weight"),
1326 vec![hidden_dim],
1327 hidden_dim,
1328 );
1329 push_bf16(
1330 &mut tensors,
1331 format!("{prefix}.self_attention_res_norm.weight"),
1332 vec![hidden_dim],
1333 hidden_dim,
1334 );
1335 push_bf16(
1336 &mut tensors,
1337 format!("{prefix}.self_attention_res_proj.weight"),
1338 vec![1, hidden_dim],
1339 hidden_dim,
1340 );
1341 push_bf16(
1342 &mut tensors,
1343 format!("{prefix}.mlp_res_norm.weight"),
1344 vec![hidden_dim],
1345 hidden_dim,
1346 );
1347 push_bf16(
1348 &mut tensors,
1349 format!("{prefix}.mlp_res_proj.weight"),
1350 vec![1, hidden_dim],
1351 hidden_dim,
1352 );
1353
1354 push_bf16(
1356 &mut tensors,
1357 format!("{prefix}.self_attn.q_a_proj.weight"),
1358 vec![q_lora_rank, hidden_dim],
1359 q_lora_rank * hidden_dim,
1360 );
1361 push_bf16(
1362 &mut tensors,
1363 format!("{prefix}.self_attn.q_a_layernorm.weight"),
1364 vec![q_lora_rank],
1365 q_lora_rank,
1366 );
1367 push_bf16(
1368 &mut tensors,
1369 format!("{prefix}.self_attn.q_b_proj.weight"),
1370 vec![num_heads * q_head_dim, q_lora_rank],
1371 num_heads * q_head_dim * q_lora_rank,
1372 );
1373 push_bf16(
1374 &mut tensors,
1375 format!("{prefix}.self_attn.kv_a_proj_with_mqa.weight"),
1376 vec![kv_lora_rank + qk_rope_head_dim, hidden_dim],
1377 (kv_lora_rank + qk_rope_head_dim) * hidden_dim,
1378 );
1379 push_bf16(
1380 &mut tensors,
1381 format!("{prefix}.self_attn.kv_a_layernorm.weight"),
1382 vec![kv_lora_rank],
1383 kv_lora_rank,
1384 );
1385 push_bf16(
1386 &mut tensors,
1387 format!("{prefix}.self_attn.kv_b_proj.weight"),
1388 vec![num_heads * (qk_nope_head_dim + v_head_dim), kv_lora_rank],
1389 num_heads * (qk_nope_head_dim + v_head_dim) * kv_lora_rank,
1390 );
1391 push_bf16(
1392 &mut tensors,
1393 format!("{prefix}.self_attn.o_proj.weight"),
1394 vec![hidden_dim, num_heads * v_head_dim],
1395 hidden_dim * num_heads * v_head_dim,
1396 );
1397 push_bf16(
1398 &mut tensors,
1399 format!("{prefix}.self_attn.g_proj.weight"),
1400 vec![num_heads * v_head_dim, hidden_dim],
1401 num_heads * v_head_dim * hidden_dim,
1402 );
1403
1404 push_bf16(
1406 &mut tensors,
1407 format!("{prefix}.block_sparse_moe.gate.weight"),
1408 vec![n_experts, hidden_dim],
1409 n_experts * hidden_dim,
1410 );
1411 let bias_bytes: Vec<u8> = vec![0.0f32; n_experts]
1412 .iter()
1413 .flat_map(|v| v.to_le_bytes())
1414 .collect();
1415 tensors.push((
1416 format!("{prefix}.block_sparse_moe.gate.e_score_correction_bias"),
1417 "F32",
1418 vec![n_experts],
1419 bias_bytes,
1420 ));
1421 push_bf16(
1422 &mut tensors,
1423 format!("{prefix}.block_sparse_moe.routed_expert_down_proj.weight"),
1424 vec![moe_hidden_dim, hidden_dim],
1425 moe_hidden_dim * hidden_dim,
1426 );
1427 push_bf16(
1428 &mut tensors,
1429 format!("{prefix}.block_sparse_moe.routed_expert_up_proj.weight"),
1430 vec![hidden_dim, moe_hidden_dim],
1431 hidden_dim * moe_hidden_dim,
1432 );
1433 push_bf16(
1434 &mut tensors,
1435 format!("{prefix}.block_sparse_moe.routed_expert_norm.weight"),
1436 vec![moe_hidden_dim],
1437 moe_hidden_dim,
1438 );
1439 push_bf16(
1440 &mut tensors,
1441 format!("{prefix}.block_sparse_moe.shared_experts.gate_proj.weight"),
1442 vec![shared_intermediate_dim, hidden_dim],
1443 shared_intermediate_dim * hidden_dim,
1444 );
1445 push_bf16(
1446 &mut tensors,
1447 format!("{prefix}.block_sparse_moe.shared_experts.down_proj.weight"),
1448 vec![hidden_dim, shared_intermediate_dim],
1449 hidden_dim * shared_intermediate_dim,
1450 );
1451 push_bf16(
1452 &mut tensors,
1453 format!("{prefix}.block_sparse_moe.shared_experts.up_proj.weight"),
1454 vec![shared_intermediate_dim, hidden_dim],
1455 shared_intermediate_dim * hidden_dim,
1456 );
1457
1458 for e in 0..n_experts {
1459 let expert_prefix = format!("{prefix}.block_sparse_moe.experts.{e}");
1460 let seed_base = (e as u32 + 1) * 10;
1461 tensors.push((
1462 format!("{expert_prefix}.w1.weight_packed"),
1463 "U8",
1464 vec![moe_intermediate_dim, moe_hidden_dim / 2],
1465 pseudo_bytes(seed_base + 1, moe_intermediate_dim * (moe_hidden_dim / 2)),
1466 ));
1467 tensors.push((
1468 format!("{expert_prefix}.w1.weight_scale"),
1469 "U8",
1470 vec![moe_intermediate_dim, moe_hidden_dim / 32],
1471 pseudo_scale_bytes(seed_base + 2, moe_intermediate_dim * (moe_hidden_dim / 32)),
1472 ));
1473 tensors.push((
1474 format!("{expert_prefix}.w2.weight_packed"),
1475 "U8",
1476 vec![moe_hidden_dim, moe_intermediate_dim / 2],
1477 pseudo_bytes(seed_base + 3, moe_hidden_dim * (moe_intermediate_dim / 2)),
1478 ));
1479 tensors.push((
1480 format!("{expert_prefix}.w2.weight_scale"),
1481 "U8",
1482 vec![moe_hidden_dim, moe_intermediate_dim / 32],
1483 pseudo_scale_bytes(seed_base + 4, moe_hidden_dim * (moe_intermediate_dim / 32)),
1484 ));
1485 tensors.push((
1486 format!("{expert_prefix}.w3.weight_packed"),
1487 "U8",
1488 vec![moe_intermediate_dim, moe_hidden_dim / 2],
1489 pseudo_bytes(seed_base + 5, moe_intermediate_dim * (moe_hidden_dim / 2)),
1490 ));
1491 tensors.push((
1492 format!("{expert_prefix}.w3.weight_scale"),
1493 "U8",
1494 vec![moe_intermediate_dim, moe_hidden_dim / 32],
1495 pseudo_scale_bytes(seed_base + 6, moe_intermediate_dim * (moe_hidden_dim / 32)),
1496 ));
1497 }
1498
1499 let shard_bytes = build_shard_owned(tensors.clone());
1500 let dir = std::env::temp_dir().join(format!(
1501 "ferrox_kimi_loader_layer_mla_moe_test_{}",
1502 std::process::id()
1503 ));
1504 std::fs::create_dir_all(&dir).unwrap();
1505 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
1506 let map_entries: Vec<String> = tensors
1507 .iter()
1508 .map(|(name, ..)| format!("\"{name}\":\"shard0.safetensors\""))
1509 .collect();
1510 let index = format!("{{\"weight_map\":{{{}}}}}", map_entries.join(","));
1511 let index_path = dir.join("model.safetensors.index.json");
1512 std::fs::write(&index_path, &index).unwrap();
1513
1514 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
1515 let mut hp = KimiRealHparams::real();
1516 hp.hidden_dim = hidden_dim;
1517 hp.mla_num_heads = num_heads;
1518 hp.mla_q_lora_rank = q_lora_rank;
1519 hp.mla_kv_lora_rank = kv_lora_rank;
1520 hp.mla_qk_nope_head_dim = qk_nope_head_dim;
1521 hp.mla_qk_rope_head_dim = qk_rope_head_dim;
1522 hp.mla_v_head_dim = v_head_dim;
1523 hp.moe_hidden_dim = moe_hidden_dim;
1524 hp.moe_intermediate_dim = moe_intermediate_dim;
1525 hp.n_experts = n_experts;
1526 hp.num_shared_experts = num_shared_experts;
1527
1528 let layer = load_kimi_layer(&shard, &hp, LayerAttentionKind::KimiMla, false, layer_idx)
1529 .expect("must load a real MLA+latent-MoE layer");
1530 std::fs::remove_dir_all(&dir).ok();
1531
1532 assert!(matches!(
1533 layer.attn,
1534 crate::kimi_decoder::KimiLayerAttention::Mla(_)
1535 ));
1536 match &layer.ffn {
1537 crate::kimi_decoder::KimiLayerFfn::Moe(moe) => {
1538 assert_eq!(moe.experts.n_experts(), n_experts);
1539 }
1540 crate::kimi_decoder::KimiLayerFfn::Dense(_) => panic!("expected Moe ffn"),
1541 }
1542 assert_eq!(layer.input_layernorm_weight.len(), hidden_dim);
1543 }
1544
1545 struct SyntheticDims {
1549 hidden_dim: usize,
1550 kda_num_heads: usize,
1551 kda_head_dim: usize,
1552 mla_num_heads: usize,
1553 mla_q_lora_rank: usize,
1554 mla_kv_lora_rank: usize,
1555 mla_qk_nope_head_dim: usize,
1556 mla_qk_rope_head_dim: usize,
1557 mla_v_head_dim: usize,
1558 dense_intermediate_dim: usize,
1559 moe_hidden_dim: usize,
1560 moe_intermediate_dim: usize,
1561 n_experts: usize,
1562 num_shared_experts: usize,
1563 }
1564
1565 #[allow(clippy::too_many_arguments)]
1572 fn push_layer_tensors(
1573 tensors: &mut Vec<(String, &'static str, Vec<usize>, Vec<u8>)>,
1574 layer_idx: usize,
1575 kind: LayerAttentionKind,
1576 is_dense: bool,
1577 d: &SyntheticDims,
1578 ) {
1579 let prefix = format!("language_model.model.layers.{layer_idx}");
1580 let push_bf16 = |tensors: &mut Vec<(String, &'static str, Vec<usize>, Vec<u8>)>,
1581 name: String,
1582 shape: Vec<usize>,
1583 n: usize| {
1584 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.05f32; n])));
1585 };
1586
1587 push_bf16(
1588 tensors,
1589 format!("{prefix}.input_layernorm.weight"),
1590 vec![d.hidden_dim],
1591 d.hidden_dim,
1592 );
1593 push_bf16(
1594 tensors,
1595 format!("{prefix}.post_attention_layernorm.weight"),
1596 vec![d.hidden_dim],
1597 d.hidden_dim,
1598 );
1599 push_bf16(
1600 tensors,
1601 format!("{prefix}.self_attention_res_norm.weight"),
1602 vec![d.hidden_dim],
1603 d.hidden_dim,
1604 );
1605 push_bf16(
1606 tensors,
1607 format!("{prefix}.self_attention_res_proj.weight"),
1608 vec![1, d.hidden_dim],
1609 d.hidden_dim,
1610 );
1611 push_bf16(
1612 tensors,
1613 format!("{prefix}.mlp_res_norm.weight"),
1614 vec![d.hidden_dim],
1615 d.hidden_dim,
1616 );
1617 push_bf16(
1618 tensors,
1619 format!("{prefix}.mlp_res_proj.weight"),
1620 vec![1, d.hidden_dim],
1621 d.hidden_dim,
1622 );
1623
1624 match kind {
1625 LayerAttentionKind::KimiKda => {
1626 let proj = d.kda_num_heads * d.kda_head_dim;
1627 for name in ["q_proj", "k_proj", "v_proj", "g_proj"] {
1628 push_bf16(
1629 tensors,
1630 format!("{prefix}.self_attn.{name}.weight"),
1631 vec![proj, d.hidden_dim],
1632 proj * d.hidden_dim,
1633 );
1634 }
1635 push_bf16(
1636 tensors,
1637 format!("{prefix}.self_attn.f_a_proj.weight"),
1638 vec![d.kda_head_dim, d.hidden_dim],
1639 d.kda_head_dim * d.hidden_dim,
1640 );
1641 push_bf16(
1642 tensors,
1643 format!("{prefix}.self_attn.f_b_proj.weight"),
1644 vec![proj, d.kda_head_dim],
1645 proj * d.kda_head_dim,
1646 );
1647 push_bf16(
1648 tensors,
1649 format!("{prefix}.self_attn.b_proj.weight"),
1650 vec![d.kda_num_heads, d.hidden_dim],
1651 d.kda_num_heads * d.hidden_dim,
1652 );
1653 push_bf16(
1654 tensors,
1655 format!("{prefix}.self_attn.o_proj.weight"),
1656 vec![d.hidden_dim, proj],
1657 d.hidden_dim * proj,
1658 );
1659 let f32_vec =
1660 |v: Vec<f32>| -> Vec<u8> { v.iter().flat_map(|x| x.to_le_bytes()).collect() };
1661 tensors.push((
1662 format!("{prefix}.self_attn.A_log"),
1663 "F32",
1664 vec![d.kda_num_heads],
1665 f32_vec(vec![0.5; d.kda_num_heads]),
1666 ));
1667 tensors.push((
1668 format!("{prefix}.self_attn.dt_bias"),
1669 "F32",
1670 vec![proj],
1671 f32_vec(vec![0.1; proj]),
1672 ));
1673 tensors.push((
1674 format!("{prefix}.self_attn.o_norm.weight"),
1675 "F32",
1676 vec![d.kda_head_dim],
1677 f32_vec(vec![1.0; d.kda_head_dim]),
1678 ));
1679 for conv_name in ["q_conv1d", "k_conv1d", "v_conv1d"] {
1680 tensors.push((
1681 format!("{prefix}.self_attn.{conv_name}.weight"),
1682 "F32",
1683 vec![proj, 1, 4],
1684 f32_vec(vec![0.1; proj * 4]),
1685 ));
1686 }
1687 }
1688 LayerAttentionKind::KimiMla => {
1689 let q_head_dim = d.mla_qk_nope_head_dim + d.mla_qk_rope_head_dim;
1690 push_bf16(
1691 tensors,
1692 format!("{prefix}.self_attn.q_a_proj.weight"),
1693 vec![d.mla_q_lora_rank, d.hidden_dim],
1694 d.mla_q_lora_rank * d.hidden_dim,
1695 );
1696 push_bf16(
1697 tensors,
1698 format!("{prefix}.self_attn.q_a_layernorm.weight"),
1699 vec![d.mla_q_lora_rank],
1700 d.mla_q_lora_rank,
1701 );
1702 push_bf16(
1703 tensors,
1704 format!("{prefix}.self_attn.q_b_proj.weight"),
1705 vec![d.mla_num_heads * q_head_dim, d.mla_q_lora_rank],
1706 d.mla_num_heads * q_head_dim * d.mla_q_lora_rank,
1707 );
1708 push_bf16(
1709 tensors,
1710 format!("{prefix}.self_attn.kv_a_proj_with_mqa.weight"),
1711 vec![d.mla_kv_lora_rank + d.mla_qk_rope_head_dim, d.hidden_dim],
1712 (d.mla_kv_lora_rank + d.mla_qk_rope_head_dim) * d.hidden_dim,
1713 );
1714 push_bf16(
1715 tensors,
1716 format!("{prefix}.self_attn.kv_a_layernorm.weight"),
1717 vec![d.mla_kv_lora_rank],
1718 d.mla_kv_lora_rank,
1719 );
1720 push_bf16(
1721 tensors,
1722 format!("{prefix}.self_attn.kv_b_proj.weight"),
1723 vec![
1724 d.mla_num_heads * (d.mla_qk_nope_head_dim + d.mla_v_head_dim),
1725 d.mla_kv_lora_rank,
1726 ],
1727 d.mla_num_heads
1728 * (d.mla_qk_nope_head_dim + d.mla_v_head_dim)
1729 * d.mla_kv_lora_rank,
1730 );
1731 push_bf16(
1732 tensors,
1733 format!("{prefix}.self_attn.o_proj.weight"),
1734 vec![d.hidden_dim, d.mla_num_heads * d.mla_v_head_dim],
1735 d.hidden_dim * d.mla_num_heads * d.mla_v_head_dim,
1736 );
1737 push_bf16(
1738 tensors,
1739 format!("{prefix}.self_attn.g_proj.weight"),
1740 vec![d.mla_num_heads * d.mla_v_head_dim, d.hidden_dim],
1741 d.mla_num_heads * d.mla_v_head_dim * d.hidden_dim,
1742 );
1743 }
1744 LayerAttentionKind::Gqa => panic!("synthetic checkpoint test never uses Gqa"),
1745 }
1746
1747 if is_dense {
1748 push_bf16(
1749 tensors,
1750 format!("{prefix}.mlp.gate_proj.weight"),
1751 vec![d.dense_intermediate_dim, d.hidden_dim],
1752 d.dense_intermediate_dim * d.hidden_dim,
1753 );
1754 push_bf16(
1755 tensors,
1756 format!("{prefix}.mlp.up_proj.weight"),
1757 vec![d.dense_intermediate_dim, d.hidden_dim],
1758 d.dense_intermediate_dim * d.hidden_dim,
1759 );
1760 push_bf16(
1761 tensors,
1762 format!("{prefix}.mlp.down_proj.weight"),
1763 vec![d.hidden_dim, d.dense_intermediate_dim],
1764 d.hidden_dim * d.dense_intermediate_dim,
1765 );
1766 } else {
1767 let shared_intermediate_dim = d.moe_intermediate_dim * d.num_shared_experts;
1768 push_bf16(
1769 tensors,
1770 format!("{prefix}.block_sparse_moe.gate.weight"),
1771 vec![d.n_experts, d.hidden_dim],
1772 d.n_experts * d.hidden_dim,
1773 );
1774 let bias_bytes: Vec<u8> = vec![0.0f32; d.n_experts]
1775 .iter()
1776 .flat_map(|v| v.to_le_bytes())
1777 .collect();
1778 tensors.push((
1779 format!("{prefix}.block_sparse_moe.gate.e_score_correction_bias"),
1780 "F32",
1781 vec![d.n_experts],
1782 bias_bytes,
1783 ));
1784 push_bf16(
1785 tensors,
1786 format!("{prefix}.block_sparse_moe.routed_expert_down_proj.weight"),
1787 vec![d.moe_hidden_dim, d.hidden_dim],
1788 d.moe_hidden_dim * d.hidden_dim,
1789 );
1790 push_bf16(
1791 tensors,
1792 format!("{prefix}.block_sparse_moe.routed_expert_up_proj.weight"),
1793 vec![d.hidden_dim, d.moe_hidden_dim],
1794 d.hidden_dim * d.moe_hidden_dim,
1795 );
1796 push_bf16(
1797 tensors,
1798 format!("{prefix}.block_sparse_moe.routed_expert_norm.weight"),
1799 vec![d.moe_hidden_dim],
1800 d.moe_hidden_dim,
1801 );
1802 push_bf16(
1803 tensors,
1804 format!("{prefix}.block_sparse_moe.shared_experts.gate_proj.weight"),
1805 vec![shared_intermediate_dim, d.hidden_dim],
1806 shared_intermediate_dim * d.hidden_dim,
1807 );
1808 push_bf16(
1809 tensors,
1810 format!("{prefix}.block_sparse_moe.shared_experts.down_proj.weight"),
1811 vec![d.hidden_dim, shared_intermediate_dim],
1812 d.hidden_dim * shared_intermediate_dim,
1813 );
1814 push_bf16(
1815 tensors,
1816 format!("{prefix}.block_sparse_moe.shared_experts.up_proj.weight"),
1817 vec![shared_intermediate_dim, d.hidden_dim],
1818 shared_intermediate_dim * d.hidden_dim,
1819 );
1820
1821 for e in 0..d.n_experts {
1822 let expert_prefix = format!("{prefix}.block_sparse_moe.experts.{e}");
1823 let seed_base = (layer_idx as u32 * 100) + (e as u32 + 1) * 10;
1824 tensors.push((
1825 format!("{expert_prefix}.w1.weight_packed"),
1826 "U8",
1827 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 2],
1828 pseudo_bytes(
1829 seed_base + 1,
1830 d.moe_intermediate_dim * (d.moe_hidden_dim / 2),
1831 ),
1832 ));
1833 tensors.push((
1834 format!("{expert_prefix}.w1.weight_scale"),
1835 "U8",
1836 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 32],
1837 pseudo_scale_bytes(
1838 seed_base + 2,
1839 d.moe_intermediate_dim * (d.moe_hidden_dim / 32),
1840 ),
1841 ));
1842 tensors.push((
1843 format!("{expert_prefix}.w2.weight_packed"),
1844 "U8",
1845 vec![d.moe_hidden_dim, d.moe_intermediate_dim / 2],
1846 pseudo_bytes(
1847 seed_base + 3,
1848 d.moe_hidden_dim * (d.moe_intermediate_dim / 2),
1849 ),
1850 ));
1851 tensors.push((
1852 format!("{expert_prefix}.w2.weight_scale"),
1853 "U8",
1854 vec![d.moe_hidden_dim, d.moe_intermediate_dim / 32],
1855 pseudo_scale_bytes(
1856 seed_base + 4,
1857 d.moe_hidden_dim * (d.moe_intermediate_dim / 32),
1858 ),
1859 ));
1860 tensors.push((
1861 format!("{expert_prefix}.w3.weight_packed"),
1862 "U8",
1863 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 2],
1864 pseudo_bytes(
1865 seed_base + 5,
1866 d.moe_intermediate_dim * (d.moe_hidden_dim / 2),
1867 ),
1868 ));
1869 tensors.push((
1870 format!("{expert_prefix}.w3.weight_scale"),
1871 "U8",
1872 vec![d.moe_intermediate_dim, d.moe_hidden_dim / 32],
1873 pseudo_scale_bytes(
1874 seed_base + 6,
1875 d.moe_intermediate_dim * (d.moe_hidden_dim / 32),
1876 ),
1877 ));
1878 }
1879 }
1880 }
1881
1882 fn build_synthetic_full_checkpoint(
1886 dir_name: &str,
1887 ) -> (
1888 std::path::PathBuf,
1889 ShardedSafetensors,
1890 crate::config::ModelConfig,
1891 KimiRealHparams,
1892 ) {
1893 let d = SyntheticDims {
1894 hidden_dim: 8,
1895 kda_num_heads: 2,
1896 kda_head_dim: 3,
1897 mla_num_heads: 1,
1898 mla_q_lora_rank: 4,
1899 mla_kv_lora_rank: 4,
1900 mla_qk_nope_head_dim: 2,
1901 mla_qk_rope_head_dim: 2,
1902 mla_v_head_dim: 2,
1903 dense_intermediate_dim: 5,
1904 moe_hidden_dim: 32,
1905 moe_intermediate_dim: 32,
1906 n_experts: 2,
1907 num_shared_experts: 1,
1908 };
1909 let vocab_size = 6;
1910
1911 let model_cfg = crate::config::ModelConfig {
1916 name: "synthetic-kimi-test",
1917 n_layers: 3,
1918 hidden_dim: d.hidden_dim,
1919 n_heads: 1,
1920 n_kv_heads: 1,
1921 head_dim: 4,
1922 vocab_size,
1923 rope_theta: 10000.0,
1924 rms_norm_eps: 1e-5,
1925 sliding_window: None,
1926 moe: ferrox_moe::MoeLayerConfig {
1927 expert_weights_scale: 1.0,
1928 n_experts: d.n_experts,
1929 n_experts_active: d.n_experts,
1930 n_shared_experts: d.num_shared_experts,
1931 hidden_dim: d.hidden_dim,
1932 expert_ffn_dim: d.moe_intermediate_dim,
1933 gating: ferrox_moe::GatingFunction::Sigmoid,
1934 norm_topk_prob: true,
1935 expert_group_count: None,
1936 expert_group_used_count: None,
1937 },
1938 n_dense_leading_layers: 1,
1939 attention: crate::config::AttentionKind::KimiHybrid(
1940 crate::config::KimiHybridAttention {
1941 kda_layers: vec![1, 2],
1942 full_attn_layers: vec![3],
1943 mla: crate::config::MlaConfig {
1944 num_heads: d.mla_num_heads,
1945 q_lora_rank: d.mla_q_lora_rank,
1946 kv_lora_rank: d.mla_kv_lora_rank,
1947 qk_nope_head_dim: d.mla_qk_nope_head_dim,
1948 qk_rope_head_dim: d.mla_qk_rope_head_dim,
1949 v_head_dim: d.mla_v_head_dim,
1950 use_output_gate: true,
1951 rope: None,
1952 },
1953 kda: crate::config::KdaConfig {
1954 num_heads: d.kda_num_heads,
1955 head_dim: d.kda_head_dim,
1956 short_conv_kernel_size: 4,
1957 gate_lower_bound: -5.0,
1958 use_full_rank_gate: true,
1959 },
1960 },
1961 ),
1962 rope_freqs: None,
1963 rope_attn_factor: 1.0,
1964 rope_dim: None,
1965 rope_freqs_long: None,
1966 rope_freqs_short: None,
1967 rope_orig_ctx: None,
1968 rope_layout: crate::config::RopeLayout::Neox,
1969 qk_norm_style: crate::capability::QkNormStyle::WholeVector,
1970 swa_pattern: None,
1971 attn_logit_softcap: None,
1972 final_logit_softcap: None,
1973 embedding_scale: None,
1974 attention_scale: None,
1975 rope_theta_swa: None,
1976 ffn_activation: crate::config::FfnActivation::Swiglu,
1977 best_effort_fields: &["synthetic test config, not a real preset"],
1978 };
1979
1980 let mut tensors: Vec<(String, &'static str, Vec<usize>, Vec<u8>)> = Vec::new();
1981 push_layer_tensors(&mut tensors, 0, LayerAttentionKind::KimiKda, true, &d);
1982 push_layer_tensors(&mut tensors, 1, LayerAttentionKind::KimiKda, false, &d);
1983 push_layer_tensors(&mut tensors, 2, LayerAttentionKind::KimiMla, false, &d);
1984
1985 let mut push_bf16_top = |name: String, shape: Vec<usize>, n: usize| {
1986 tensors.push((name, "BF16", shape, bf16_bytes(&vec![0.02f32; n])));
1987 };
1988 push_bf16_top(
1989 "language_model.model.embed_tokens.weight".to_string(),
1990 vec![vocab_size, d.hidden_dim],
1991 vocab_size * d.hidden_dim,
1992 );
1993 push_bf16_top(
1994 "language_model.lm_head.weight".to_string(),
1995 vec![vocab_size, d.hidden_dim],
1996 vocab_size * d.hidden_dim,
1997 );
1998 push_bf16_top(
1999 "language_model.model.norm.weight".to_string(),
2000 vec![d.hidden_dim],
2001 d.hidden_dim,
2002 );
2003 push_bf16_top(
2004 "language_model.model.output_attn_res_norm.weight".to_string(),
2005 vec![d.hidden_dim],
2006 d.hidden_dim,
2007 );
2008 push_bf16_top(
2009 "language_model.model.output_attn_res_proj.weight".to_string(),
2010 vec![1, d.hidden_dim],
2011 d.hidden_dim,
2012 );
2013
2014 let shard_bytes = build_shard_owned(tensors.clone());
2015 let dir = std::env::temp_dir().join(dir_name);
2016 std::fs::create_dir_all(&dir).unwrap();
2017 std::fs::write(dir.join("shard0.safetensors"), &shard_bytes).unwrap();
2018 let map_entries: Vec<String> = tensors
2019 .iter()
2020 .map(|(name, ..)| format!("\"{name}\":\"shard0.safetensors\""))
2021 .collect();
2022 let index = format!("{{\"weight_map\":{{{}}}}}", map_entries.join(","));
2023 let index_path = dir.join("model.safetensors.index.json");
2024 std::fs::write(&index_path, &index).unwrap();
2025
2026 let shard = ShardedSafetensors::open_index(&index_path).expect("must open index");
2027 let hp = KimiRealHparams {
2028 hidden_dim: d.hidden_dim,
2029 kda_num_heads: d.kda_num_heads,
2030 kda_head_dim: d.kda_head_dim,
2031 mla_num_heads: d.mla_num_heads,
2032 mla_q_lora_rank: d.mla_q_lora_rank,
2033 mla_kv_lora_rank: d.mla_kv_lora_rank,
2034 mla_qk_nope_head_dim: d.mla_qk_nope_head_dim,
2035 mla_qk_rope_head_dim: d.mla_qk_rope_head_dim,
2036 mla_v_head_dim: d.mla_v_head_dim,
2037 dense_intermediate_dim: d.dense_intermediate_dim,
2038 moe_hidden_dim: d.moe_hidden_dim,
2039 moe_intermediate_dim: d.moe_intermediate_dim,
2040 n_experts: d.n_experts,
2041 num_shared_experts: d.num_shared_experts,
2042 };
2043
2044 (dir, shard, model_cfg, hp)
2045 }
2046
2047 #[test]
2053 fn store_backed_kimi_experts_produce_bit_identical_outputs() {
2054 let (dir, shard, model_cfg, hp) =
2055 build_synthetic_full_checkpoint("ferrox_kimi_store_equivalence_test");
2056
2057 let eager = load_kimi_checkpoint(&shard, &model_cfg, &hp).expect("eager load");
2058 let mla_cfg = crate::config::MlaConfig {
2059 num_heads: hp.mla_num_heads,
2060 q_lora_rank: hp.mla_q_lora_rank,
2061 kv_lora_rank: hp.mla_kv_lora_rank,
2062 qk_nope_head_dim: hp.mla_qk_nope_head_dim,
2063 qk_rope_head_dim: hp.mla_qk_rope_head_dim,
2064 v_head_dim: hp.mla_v_head_dim,
2065 use_output_gate: true,
2066 rope: None,
2067 };
2068 let kda_cfg = crate::config::KdaConfig {
2069 num_heads: hp.kda_num_heads,
2070 head_dim: hp.kda_head_dim,
2071 short_conv_kernel_size: 4,
2072 gate_lower_bound: -5.0,
2073 use_full_rank_gate: true,
2074 };
2075 let dec_cfg = crate::kimi_decoder::KimiDecoderConfig {
2076 attn_res_block_size: 12,
2077 rms_norm_eps: 1e-5,
2078 situ_beta: 4.0,
2079 situ_linear_beta: 25.0,
2080 moe: crate::latent_moe::KimiMoeConfig {
2081 n_experts_active: hp.n_experts,
2082 moe_renormalize: true,
2083 routed_scaling_factor: 1.0,
2084 situ_beta: 4.0,
2085 situ_linear_beta: 25.0,
2086 rms_norm_eps: 1e-5,
2087 },
2088 };
2089
2090 for budget in [64 * 1024 * 1024u64, 1u64] {
2091 let stored =
2092 load_kimi_checkpoint_with_expert_cache(&shard, &model_cfg, &hp, Some(budget))
2093 .expect("store-backed load");
2094 let mut state_a = crate::kimi_decoder::KimiDecodeState::new(&eager, &kda_cfg);
2095 let mut state_b = crate::kimi_decoder::KimiDecodeState::new(&stored, &kda_cfg);
2096 for &tok in &[1usize, 3, 0, 2] {
2097 let a = crate::kimi_decoder::kimi_forward_token(
2098 &eager,
2099 &dec_cfg,
2100 &mla_cfg,
2101 &kda_cfg,
2102 tok,
2103 &mut state_a,
2104 );
2105 let b = crate::kimi_decoder::kimi_forward_token(
2106 &stored,
2107 &dec_cfg,
2108 &mla_cfg,
2109 &kda_cfg,
2110 tok,
2111 &mut state_b,
2112 );
2113 assert_eq!(
2114 a, b,
2115 "budget={budget}: store-backed Kimi output must be bit-identical"
2116 );
2117 }
2118 }
2119 std::fs::remove_dir_all(&dir).ok();
2120 }
2121
2122 #[test]
2123 fn load_kimi_checkpoint_assembles_every_real_layer_kind() {
2124 let (dir, shard, model_cfg, hp) =
2125 build_synthetic_full_checkpoint("ferrox_kimi_loader_full_checkpoint_test");
2126
2127 let weights = load_kimi_checkpoint(&shard, &model_cfg, &hp)
2128 .expect("must assemble a complete synthetic checkpoint");
2129 std::fs::remove_dir_all(&dir).ok();
2130
2131 assert_eq!(weights.layers.len(), 3);
2132 assert!(matches!(
2133 weights.layers[0].ffn,
2134 crate::kimi_decoder::KimiLayerFfn::Dense(_)
2135 ));
2136 assert!(matches!(
2137 weights.layers[0].attn,
2138 crate::kimi_decoder::KimiLayerAttention::Kda(_)
2139 ));
2140 assert!(matches!(
2141 weights.layers[1].ffn,
2142 crate::kimi_decoder::KimiLayerFfn::Moe(_)
2143 ));
2144 assert!(matches!(
2145 weights.layers[1].attn,
2146 crate::kimi_decoder::KimiLayerAttention::Kda(_)
2147 ));
2148 assert!(matches!(
2149 weights.layers[2].ffn,
2150 crate::kimi_decoder::KimiLayerFfn::Moe(_)
2151 ));
2152 assert!(matches!(
2153 weights.layers[2].attn,
2154 crate::kimi_decoder::KimiLayerAttention::Mla(_)
2155 ));
2156 assert_eq!(weights.embedding.rows(), model_cfg.vocab_size);
2157 assert_eq!(weights.embedding.cols(), hp.hidden_dim);
2158 assert_eq!(weights.output_head.rows(), model_cfg.vocab_size);
2159 assert_eq!(weights.final_norm_weight.len(), hp.hidden_dim);
2160
2161 let mla_cfg = crate::config::MlaConfig {
2165 num_heads: hp.mla_num_heads,
2166 q_lora_rank: hp.mla_q_lora_rank,
2167 kv_lora_rank: hp.mla_kv_lora_rank,
2168 qk_nope_head_dim: hp.mla_qk_nope_head_dim,
2169 qk_rope_head_dim: hp.mla_qk_rope_head_dim,
2170 v_head_dim: hp.mla_v_head_dim,
2171 use_output_gate: true,
2172 rope: None,
2173 };
2174 let kda_cfg = crate::config::KdaConfig {
2175 num_heads: hp.kda_num_heads,
2176 head_dim: hp.kda_head_dim,
2177 short_conv_kernel_size: 4,
2178 gate_lower_bound: -5.0,
2179 use_full_rank_gate: true,
2180 };
2181 let decoder_cfg = crate::kimi_decoder::KimiDecoderConfig {
2182 attn_res_block_size: 12,
2183 rms_norm_eps: 1e-5,
2184 situ_beta: 4.0,
2185 situ_linear_beta: 25.0,
2186 moe: crate::latent_moe::KimiMoeConfig {
2187 n_experts_active: hp.n_experts,
2188 moe_renormalize: true,
2189 routed_scaling_factor: 1.0,
2190 situ_beta: 4.0,
2191 situ_linear_beta: 25.0,
2192 rms_norm_eps: 1e-5,
2193 },
2194 };
2195 let mut state = crate::kimi_decoder::KimiDecodeState::new(&weights, &kda_cfg);
2196 let logits = crate::kimi_decoder::kimi_forward_token(
2197 &weights,
2198 &decoder_cfg,
2199 &mla_cfg,
2200 &kda_cfg,
2201 0,
2202 &mut state,
2203 );
2204 assert_eq!(logits.len(), model_cfg.vocab_size);
2205 assert!(logits.iter().all(|v| v.is_finite()));
2206 }
2207}