Skip to main content

frink_models/
kimi_loader.rs

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