1use anyhow::{Context, Result, bail};
13
14use super::kda::{KdaConfig, KdaState};
15use super::layer::{K3Graph, MixerKind};
16
17#[derive(Clone, Debug, Default)]
19pub struct MlaKv {
20 pub k: Vec<f32>,
22 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#[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 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 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}