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