Skip to main content

ferrox_models/
kimi_loader.rs

1//! Loads `ferrox-models::kimi_decoder` weights from Kimi K3's real
2//! safetensors checkpoint (via `ferrox-safetensors::ShardedSafetensors`),
3//! using the exact real tensor names/shapes/dtypes fetched live from
4//! `huggingface.co/moonshotai/Kimi-K3` (a real shard header, not
5//! guessed) -- confirmed to match this crate's `kda`/`mla`/`latent_moe`/
6//! `kimi_decoder` struct field names and shapes exactly.
7//!
8//! One real, non-obvious fact confirmed by reading actual tensor shapes
9//! rather than assuming they match `modeling_kimi_linear.py`'s
10//! `KimiDeltaAttention.__init__` literally: `self_attn.A_log`'s real
11//! on-disk shape is `[128]`, not `[num_heads]` = `[96]` (confirmed
12//! independently by `self_attn.b_proj.weight`'s real shape `[96,
13//! 7168]`, which is unambiguously `[num_heads, hidden_dim]`). The real
14//! `fused_recurrent_kda` kernel only ever indexes `A_log[i_hv]` for
15//! `i_hv` in `0..num_heads`, so this is real, harmless padding (likely
16//! to a GPU-friendly round size) rather than a spec mismatch -- this
17//! loader reads the real 128-element tensor but only uses its first
18//! `num_heads` elements, matching what the real kernel actually
19//! consumes.
20//!
21//! Dequantizes small `F32`/`BF16` tensors (per-head attention
22//! parameters, layernorms, dense-FFN and shared-expert projections) to
23//! owned `f32` eagerly at load time -- matching this project's
24//! established BF16-handling convention, e.g. `ferrox-models::loader`'s
25//! GGUF path -- since these are cheap regardless. Routed-expert `MXFP4`
26//! weights are the one format this loader does **not** eagerly
27//! dequantize: `load_mxfp4_weight_matrix` builds a zero-copy
28//! `WeightMatrix::Mxfp4` (mmap-backed `packed`/`scale` buffers, real
29//! Kimi K3 stores these as two separate tensors per projection -- see
30//! `ferrox_quant::dot_mxfp4_row_f32`'s doc comment), matching the
31//! zero-copy-mmap-plus-fused-dot discipline every other quantized
32//! format in this codebase already uses. This is a real fix, not just a
33//! design preference: real-hardware testing (rented 62GB instance, see
34//! docs/MODELS.md) found that eagerly dequantizing all 896 of a real MoE
35//! layer's routed experts to owned `f32` needs roughly 117GB of RAM and
36//! reproducibly OOM-killed a 62GB rented instance -- the zero-copy path
37//! here keeps a loaded layer's resident memory close to its on-disk
38//! size (the real MXFP4 packing ratio: 2 values/byte plus one scale
39//! byte per 32 values) instead of expanding every value to 4 bytes
40//! whether it's ever used by real top-k routing or not.
41
42use 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
64/// Reads any real tensor as an owned `f32` vector, dispatching on its
65/// real declared dtype (`F32` direct, `BF16` dequantized) -- exposed
66/// `pub` since not every real weight (e.g. the per-layer
67/// `input_layernorm.weight`/`post_attention_layernorm.weight`, which
68/// aren't nested inside `KdaAttnWeights`/`MlaAttnWeights`/
69/// `DenseMlpWeights`/`BlockResidualWeights`) has a dedicated loader
70/// function above.
71pub 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
106/// Loads one KDA-attention layer's weights (real tensor names under
107/// `{prefix}.self_attn.*`). `num_heads`/`head_dim`/`hidden_dim` must
108/// match the real config (Kimi K3: 96/128/7168) -- passed explicitly
109/// rather than hardcoded so this loader can also be exercised against
110/// small synthetic on-disk fixtures in tests.
111pub 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    // Real padding: only the first `num_heads` of A_log's real on-disk
121    // elements are ever read by the real kernel -- see module doc
122    // comment.
123    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        // Real on-disk shape is [projection_size, 1, kernel_size]; the
145        // middle dim is always 1 (depthwise conv), so the raw bytes are
146        // already exactly [projection_size, kernel_size] flattened --
147        // no reshape needed, just read as a flat vec.
148        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/// Loads one Gated-MLA-attention layer's weights (real tensor names
188/// under `{prefix}.self_attn.*`).
189#[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
244/// Loads the dense leading layer's feed-forward block (real tensor
245/// names under `{prefix}.mlp.*`) -- Kimi K3's layer 0 only
246/// (`first_k_dense_replace`=1).
247pub 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
275/// The four block-residual weight vectors real Kimi K3 attaches to
276/// every layer (`{prefix}.self_attention_res_{norm,proj}.weight`,
277/// `{prefix}.mlp_res_{norm,proj}.weight`). Real on-disk `*_proj.weight`
278/// shape is `[1, hidden_dim]` (a `Linear(hidden_dim, 1)`'s weight); the
279/// raw bytes are already exactly `[hidden_dim]` flattened.
280pub 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
305/// Loads one MXFP4-quantized weight matrix from its real two-separate-
306/// tensors shape (`*.weight_packed` + `*.weight_scale`, confirmed
307/// against a real shard header -- see `ferrox_quant::dot_mxfp4_row_f32`'s
308/// doc comment) as a zero-copy `WeightMatrix::Mxfp4`: both buffers are
309/// mmap-backed views (`ShardedSafetensors::tensor_mapped_range`), never
310/// copied into an owned buffer, let alone dequantized to `f32` --
311/// `apply`/`apply_batch` dispatch straight to the fused
312/// `dot_mxfp4_row_f32` kernel. See this module's doc comment for why
313/// this matters at real scale (a real, measured ~117GB RAM difference
314/// for one full 896-expert MoE layer).
315fn 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/// Loads one routed expert's real MXFP4 weights
354/// (`{prefix}.experts.{expert_idx}.{w1,w2,w3}.{weight_packed,weight_scale}`).
355/// `moe_hidden_dim` is the real *latent* dimension
356/// (`routed_expert_hidden_size`=3584 for Kimi K3, not the outer
357/// `hidden_dim`=7168 -- see `ferrox-models::latent_moe`'s module doc
358/// comment); `moe_intermediate_dim` is the per-expert FFN size (3072).
359/// Per-layer byte layout of one store-backed Kimi routed expert's
360/// combined buffer: w1_packed, w1_scale, w2_packed, w2_scale,
361/// w3_packed, w3_scale concatenated in that fixed order. Every expert
362/// in a Kimi layer has identical dims, so one layout serves the layer.
363#[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, // w1 packed: [m, h]
375            m * h / g, // w1 scale
376            h * m / 2, // w2 packed: [h, m]
377            h * m / g, // w2 scale
378            m * h / 2, // w3 packed: [m, h]
379            m * h / g, // w3 scale
380        ]
381    }
382
383    pub fn total_bytes(&self) -> usize {
384        self.seg_lens().iter().sum()
385    }
386
387    /// Builds temporary two-buffer MXFP4 `WeightMatrix` views over a
388    /// leased buffer; each view's `WeightBytes::Shared` clone keeps
389    /// the cache entry pinned for the view's lifetime.
390    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
415/// [`ExpertSource`] over a Kimi safetensors checkpoint: each expert's
416/// six tensors (three matrices' packed+scale buffers) are read
417/// positionally from the owning shard files and concatenated in
418/// `KimiStoredExpertLayout`'s fixed order.
419pub struct KimiExpertSource {
420    files: Vec<std::fs::File>,
421    /// (layer, expert) -> six (file index, offset, len) segments.
422    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/// Loads one full MoE layer: the gate (with its real aux-loss-free
494/// `e_score_correction_bias`), the shared down/up latent projections +
495/// norm, every routed expert (`n_experts`, real MXFP4), and the shared
496/// expert (real `BF16`, on the full `hidden_dim`, not the latent space
497/// -- see `ferrox-models::latent_moe`'s module doc comment).
498#[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
572/// Kimi K3's real per-layer hyperparameters needed to load any layer
573/// (not tied to `ferrox_moe::MoeLayerConfig`/`ferrox_models::ModelConfig`,
574/// neither of which model the "latent MoE" down-projected dimension or
575/// the dense leading layer's own intermediate size -- kept as a small,
576/// dedicated struct here rather than widening those shared types for
577/// one model's real values).
578pub 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    /// Kimi K3's real published values (`config.json`).
597    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
617/// Loads any one layer (KDA or Gated-MLA attention, dense or MoE FFN),
618/// dispatching on `kind`/`is_dense` -- pass
619/// `ModelConfig::layer_attention_kind(layer_idx)`/
620/// `ModelConfig::layer_is_dense(layer_idx)` for Kimi K3's real per-layer
621/// topology.
622pub 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
695/// Loads a complete `KimiDecoderWeights` -- every one of `model_cfg`'s
696/// real layers (dispatched per-layer via `model_cfg.layer_attention_kind`/
697/// `layer_is_dense`, driven by `hp`'s per-layer dimensions), plus the
698/// real top-level tensors (real names confirmed against a real shard
699/// header: `language_model.model.embed_tokens.weight`,
700/// `language_model.lm_head.weight`, `language_model.model.norm.weight`,
701/// `language_model.model.output_attn_res_{norm,proj}.weight`). This is
702/// the assembly step `load_kimi_layer` itself doesn't do -- calling it
703/// once per real layer and building the surrounding `KimiDecoderWeights`
704/// -- analogous to `ferrox-models::loader::Decoder::from_gguf`, but for
705/// Kimi K3's real safetensors format. Not blocked on anything (the
706/// zero-copy MXFP4 fix removes the memory obstacle a full loader would
707/// otherwise hit for every non-dense layer's routed experts); simply
708/// not runnable against the real 2.8T-parameter checkpoint in this
709/// environment (96 shards, 1.56TB) -- tested here against small
710/// synthetic on-disk fixtures instead, real safetensors bytes and real
711/// tensor names throughout.
712/// Like [`load_kimi_checkpoint`], but with `expert_cache_bytes:
713/// Some(budget)` every MoE layer's routed experts are converted to
714/// store-backed lazy materialization after loading: one bounded,
715/// lease-protected `ExpertStore` shared by the whole model reads each
716/// expert's six tensors positionally from the owning shard files on
717/// miss, instead of holding 896 expert objects per layer resident.
718/// Attention, dense layers, shared experts, router/projections,
719/// embeddings, and the output head are untouched. Bit-identical to
720/// the eager path (same bytes, same kernels) -- pinned by the
721/// equivalence test against the synthetic multi-layer checkpoint.
722pub 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    // Collect every MoE layer's per-expert file segments, then swap
734    // each layer's backing to the one shared store.
735    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(); // (layer_idx, n_experts)
745
746    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    /// Builds a real on-disk safetensors file (header + raw bytes)
857    /// containing exactly one BF16-widened-from-f32 tensor for the
858    /// given name/shape -- BF16 chosen since it's the dominant real
859    /// dtype in Kimi K3's checkpoint, exercising the dequant path.
860    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        // Real on-disk A_log has 128 elements but only num_heads(=2
961        // here) are ever consumed -- confirm the loader truncates
962        // rather than asserting a shape match against the full tensor.
963        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    /// Deterministic byte generator (no external `rand` dependency in
987    /// this crate) -- any byte pattern is a structurally valid MXFP4
988    /// block (dequant doesn't depend on the bytes coming from a real
989    /// quantizer), so this just needs to be varied, not random.
990    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    /// Like `pseudo_bytes`, but clamped to a realistic E8M0 scale range
1001    /// (roughly `2^-127` to `2^53`). Scale bytes near the top of the
1002    /// real `u8` range (255 reserved for NaN by the OCP spec and
1003    /// deliberately not special-cased by `ferrox_quant::dequant_mxfp4_row`,
1004    /// matching real `ggml_e8m0_to_fp32`'s own documented limitation; and
1005    /// bytes up to ~252 combined with E2M1's max magnitude of 6 can
1006    /// legitimately overflow `f32::MAX`) are real OCP MX behavior, not a
1007    /// bug -- just not representative of any real *trained* weight's
1008    /// scale, and not what this test (confirming the loader wires real
1009    /// bytes through correctly) is checking for.
1010    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        // Smallest valid dims: MXFP4_GROUP_SIZE=32, so every in_dim here
1020        // must be a multiple of 32.
1021        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    /// Like `build_shard`, but takes owned `String`/`Vec<usize>` tensor
1113    /// descriptors so callers can build the list programmatically
1114    /// (needed for `load_kimi_layer`'s tests, which have far more
1115    /// tensors than the hand-written fixtures above).
1116    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        // MLA attention tensors.
1355        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        // Latent-MoE tensors.
1405        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    /// Dims for the small synthetic checkpoint
1546    /// `load_kimi_checkpoint_assembles_every_real_layer_kind` builds --
1547    /// every field mirrors `KimiRealHparams`, just at test scale.
1548    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    /// Appends one real layer's tensor set (KDA or MLA attention, dense
1566    /// or MoE FFN, per `kind`/`is_dense`) to `tensors`, matching the
1567    /// exact real tensor names/shapes `load_kimi_layer` expects --
1568    /// shared by `load_kimi_checkpoint_assembles_every_real_layer_kind`
1569    /// across all 3 of its synthetic layers to avoid repeating each
1570    /// layer's ~15-30 tensor descriptors by hand.
1571    #[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    /// Builds the 3-layer synthetic checkpoint (dense+KDA, MoE+KDA,
1883    /// MoE+MLA) on disk and opens it -- shared by the assembly test and
1884    /// the store-backed equivalence test. Caller removes `dir`.
1885    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        // 3 real layers: 0 = dense+KDA (matches Kimi K3's real layer 0),
1912        // 1 = MoE+KDA, 2 = MoE+MLA -- covering every real
1913        // attention/FFN combination `load_kimi_checkpoint` must
1914        // dispatch correctly.
1915        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    /// The unchanged-output gate for Kimi expert streaming: the same
2048    /// synthetic checkpoint loaded eagerly vs. store-backed (generous
2049    /// AND smaller-than-one-expert budgets) must produce bit-identical
2050    /// forward-pass outputs -- same bytes, same kernels, assert_eq on
2051    /// f32 vectors with no tolerance.
2052    #[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        // Run a real forward pass through the fully-assembled checkpoint
2162        // to confirm every piece composes correctly end to end, not
2163        // just that each layer loads.
2164        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}