atlas_core/kimi_k3/
cache.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Hybrid cache: paged MLA KV + KDA recurrent/conv state.
4//!
5//! Prefix-cache restore reuses C3 CPU semantics: snapshot/restore the
6//! per-layer [`LayerCache`] (KDA conv/recurrent or MLA KV).
7//!
8//! Host bytes are the prefix snapshot. Decode hot path must use
9//! `spark_model::kimi_k3::DeviceHybridCache` (device-resident conv/recurrent
10//! and MLA KV). Per-token D2H/H2D of those buffers is not this cache.
11
12use anyhow::{Context, Result, bail};
13
14use super::kda::{KdaConfig, KdaState};
15use super::layer::{K3Graph, MixerKind};
16
17/// One MLA layer's host KV (unpaged CPU stand-in).
18#[derive(Clone, Debug, Default)]
19pub struct MlaKv {
20    /// Packed keys `[T, H, dq]`.
21    pub k: Vec<f32>,
22    /// Packed values `[T, H, dv]`.
23    pub v: Vec<f32>,
24    pub seq_len: usize,
25}
26
27#[derive(Clone, Debug)]
28pub enum LayerCache {
29    Kda(KdaState),
30    Mla(MlaKv),
31}
32
33/// Per-sequence hybrid cache. Slot identity is the layer index; a prefix
34/// hit that writes KDA state into the wrong slot is the C4 mutant.
35#[derive(Clone, Debug)]
36pub struct HybridCache {
37    pub layers: Vec<LayerCache>,
38}
39
40impl HybridCache {
41    pub fn from_graph(graph: &K3Graph, kda: &KdaConfig) -> Self {
42        let layers = graph
43            .layers
44            .iter()
45            .map(|l| match l.mixer {
46                MixerKind::Kda => LayerCache::Kda(KdaState::new(kda)),
47                MixerKind::Mla => LayerCache::Mla(MlaKv::default()),
48            })
49            .collect();
50        Self { layers }
51    }
52
53    pub fn kda_mut(&mut self, layer: usize) -> Option<&mut KdaState> {
54        match self.layers.get_mut(layer) {
55            Some(LayerCache::Kda(s)) => Some(s),
56            _ => None,
57        }
58    }
59
60    pub fn mla_mut(&mut self, layer: usize) -> Option<&mut MlaKv> {
61        match self.layers.get_mut(layer) {
62            Some(LayerCache::Mla(s)) => Some(s),
63            _ => None,
64        }
65    }
66}
67
68impl MlaKv {
69    /// Append one token's packed K/V (`[H, dq]` / `[H, dv]`).
70    pub fn append(&mut self, k: &[f32], v: &[f32]) {
71        self.k.extend_from_slice(k);
72        self.v.extend_from_slice(v);
73        self.seq_len += 1;
74    }
75}
76
77impl LayerCache {
78    /// Host blob for Marconi aux / C3 prefix-cache restore.
79    pub fn to_bytes(&self) -> Vec<u8> {
80        let mut b = Vec::new();
81        match self {
82            LayerCache::Kda(s) => {
83                b.push(1);
84                push_f32s(&mut b, &s.conv);
85                push_f32s(&mut b, &s.recurrent);
86            }
87            LayerCache::Mla(kv) => {
88                b.push(2);
89                b.extend_from_slice(&(kv.seq_len as u32).to_le_bytes());
90                push_f32s(&mut b, &kv.k);
91                push_f32s(&mut b, &kv.v);
92            }
93        }
94        b
95    }
96
97    pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
98        let tag = bytes.first().copied().context("empty K3 LayerCache blob")?;
99        let rest = &bytes[1..];
100        match tag {
101            1 => {
102                let (conv, rest) = take_f32s(rest)?;
103                let (recurrent, rest) = take_f32s(rest)?;
104                if !rest.is_empty() {
105                    bail!("KDA LayerCache blob has trailing bytes");
106                }
107                Ok(LayerCache::Kda(KdaState { conv, recurrent }))
108            }
109            2 => {
110                if rest.len() < 4 {
111                    bail!("MLA LayerCache blob truncated seq_len");
112                }
113                let seq_len = u32::from_le_bytes(rest[..4].try_into().unwrap()) as usize;
114                let (k, rest) = take_f32s(&rest[4..])?;
115                let (v, rest) = take_f32s(rest)?;
116                if !rest.is_empty() {
117                    bail!("MLA LayerCache blob has trailing bytes");
118                }
119                Ok(LayerCache::Mla(MlaKv { k, v, seq_len }))
120            }
121            t => bail!("unknown K3 LayerCache tag {t}"),
122        }
123    }
124}
125
126fn push_f32s(b: &mut Vec<u8>, xs: &[f32]) {
127    b.extend_from_slice(&(xs.len() as u32).to_le_bytes());
128    for x in xs {
129        b.extend_from_slice(&x.to_le_bytes());
130    }
131}
132
133fn take_f32s(bytes: &[u8]) -> Result<(Vec<f32>, &[u8])> {
134    if bytes.len() < 4 {
135        bail!("LayerCache f32 vec truncated length");
136    }
137    let n = u32::from_le_bytes(bytes[..4].try_into().unwrap()) as usize;
138    let need = 4 + n.checked_mul(4).context("LayerCache f32 overflow")?;
139    if bytes.len() < need {
140        bail!("LayerCache f32 vec truncated body");
141    }
142    let mut v = Vec::with_capacity(n);
143    for chunk in bytes[4..need].chunks_exact(4) {
144        v.push(f32::from_le_bytes(chunk.try_into().unwrap()));
145    }
146    Ok((v, &bytes[need..]))
147}
148
149#[cfg(test)]
150mod tests {
151    use super::*;
152    use crate::config::parse_config;
153    use crate::kimi_k3::layer::K3Graph;
154
155    #[test]
156    fn twin_cache_slots_follow_mixer() {
157        const TWIN: &str = include_str!("../../../../docs/k3/fixtures/Kimi-K3-0.40B-config.json");
158        let c = parse_config(TWIN).unwrap();
159        let g = K3Graph::from_config(&c);
160        let kda = KdaConfig {
161            heads: 8,
162            head_dim: 32,
163            conv_kernel: 4,
164            gate_lower_bound: Some(-5.0),
165            use_full_rank_gate: true,
166        };
167        let cache = HybridCache::from_graph(&g, &kda);
168        assert_eq!(cache.layers.len(), 8);
169        for i in [0, 1, 2, 4, 5, 6] {
170            assert!(matches!(cache.layers[i], LayerCache::Kda(_)));
171        }
172        for i in [3, 7] {
173            assert!(matches!(cache.layers[i], LayerCache::Mla(_)));
174        }
175    }
176
177    #[test]
178    fn layer_cache_bytes_roundtrip_kda_and_mla() {
179        let kda = KdaConfig {
180            heads: 1,
181            head_dim: 2,
182            conv_kernel: 4,
183            gate_lower_bound: Some(-5.0),
184            use_full_rank_gate: true,
185        };
186        let mut k = LayerCache::Kda(KdaState::new(&kda));
187        if let LayerCache::Kda(s) = &mut k {
188            s.conv[0] = 1.25;
189            s.recurrent[0] = -0.5;
190        }
191        let back = LayerCache::from_bytes(&k.to_bytes()).unwrap();
192        match back {
193            LayerCache::Kda(s) => {
194                assert_eq!(s.conv[0], 1.25);
195                assert_eq!(s.recurrent[0], -0.5);
196            }
197            LayerCache::Mla(_) => panic!("KDA roundtrip"),
198        }
199        let mut m = LayerCache::Mla(MlaKv::default());
200        if let LayerCache::Mla(kv) = &mut m {
201            kv.append(&[1.0, 2.0], &[3.0, 4.0]);
202        }
203        let back = LayerCache::from_bytes(&m.to_bytes()).unwrap();
204        match back {
205            LayerCache::Mla(kv) => {
206                assert_eq!(kv.seq_len, 1);
207                assert_eq!(kv.k, vec![1.0, 2.0]);
208                assert_eq!(kv.v, vec![3.0, 4.0]);
209            }
210            LayerCache::Kda(_) => panic!("MLA roundtrip"),
211        }
212    }
213
214    #[test]
215    fn trash_kda_state_after_prefix_clone_diverges() {
216        let kda = KdaConfig::twin_0_40b();
217        let mut cache = HybridCache {
218            layers: vec![LayerCache::Kda(KdaState::new(&kda))],
219        };
220        if let LayerCache::Kda(s) = &mut cache.layers[0] {
221            s.conv[0] = 0.3;
222            s.recurrent[1] = -0.2;
223        }
224        let prefix = cache.clone();
225        if let LayerCache::Kda(s) = &mut cache.layers[0] {
226            for x in &mut s.conv {
227                *x = 7.0;
228            }
229            for x in &mut s.recurrent {
230                *x = 7.0;
231            }
232        }
233        assert_ne!(
234            cache.layers[0].to_bytes(),
235            prefix.layers[0].to_bytes(),
236            "RST known-bad: trash-all KDA conv+recurrent after prefix clone must change the blob"
237        );
238    }
239
240    #[test]
241    fn wrong_mla_kv_row_after_append_diverges() {
242        let mut kv = MlaKv::default();
243        kv.append(&[1.0, 0.0], &[0.0, 1.0]);
244        kv.append(&[2.0, 0.0], &[0.0, 2.0]);
245        let clean = kv.clone();
246        let k_stride = kv.k.len() / kv.seq_len;
247        let v_stride = kv.v.len() / kv.seq_len;
248        let last = kv.seq_len - 1;
249        for i in 0..k_stride {
250            kv.k.swap(i, last * k_stride + i);
251        }
252        for i in 0..v_stride {
253            kv.v.swap(i, last * v_stride + i);
254        }
255        assert_ne!(
256            kv.k, clean.k,
257            "RST known-bad: swapped MLA K row must diverge"
258        );
259        assert_ne!(
260            kv.v, clean.v,
261            "RST known-bad: swapped MLA V row must diverge"
262        );
263    }
264}