1use anyhow::Result;
21use crate::tensor_core::Tensor;
22
23pub const DEFAULT_MAX_SEQ_LEN: usize = 2048;
25
26#[derive(Debug, Clone)]
28pub struct LayerKVCache {
29 k_buffer: Option<candle_core::Tensor>,
31 v_buffer: Option<candle_core::Tensor>,
33 current_len: usize,
35 max_seq_len: usize,
37}
38
39impl LayerKVCache {
40 pub fn new() -> Self {
42 Self::with_capacity(DEFAULT_MAX_SEQ_LEN)
43 }
44
45 pub fn with_capacity(max_seq_len: usize) -> Self {
47 Self {
48 k_buffer: None,
49 v_buffer: None,
50 current_len: 0,
51 max_seq_len,
52 }
53 }
54
55 pub fn is_empty(&self) -> bool {
57 self.current_len == 0
58 }
59
60 pub fn clear(&mut self) {
62 self.current_len = 0;
63 }
65
66 pub fn deallocate(&mut self) {
68 self.k_buffer = None;
69 self.v_buffer = None;
70 self.current_len = 0;
71 }
72
73 pub fn seq_len(&self) -> usize {
75 self.current_len
76 }
77
78 pub fn max_seq_len(&self) -> usize {
80 self.max_seq_len
81 }
82
83 pub fn append(&mut self, new_k: &candle_core::Tensor, new_v: &candle_core::Tensor) -> Result<()> {
92 let new_len = new_k.dims()[2];
93
94 if self.k_buffer.is_none() {
96 let shape = new_k.dims();
97 let buffer_shape = (shape[0], shape[1], self.max_seq_len, shape[3]);
99 self.k_buffer = Some(candle_core::Tensor::zeros(
100 buffer_shape,
101 new_k.dtype(),
102 new_k.device(),
103 )?);
104 self.v_buffer = Some(candle_core::Tensor::zeros(
105 buffer_shape,
106 new_v.dtype(),
107 new_v.device(),
108 )?);
109 }
110
111 if self.current_len + new_len > self.max_seq_len {
113 anyhow::bail!(
114 "KV cache overflow: current_len={}, new_len={}, max_seq_len={}",
115 self.current_len,
116 new_len,
117 self.max_seq_len
118 );
119 }
120
121 self.k_buffer
124 .as_mut()
125 .unwrap()
126 .slice_set(new_k, 2, self.current_len)?;
127 self.v_buffer
128 .as_mut()
129 .unwrap()
130 .slice_set(new_v, 2, self.current_len)?;
131
132 self.current_len += new_len;
133 Ok(())
134 }
135
136 pub fn get_kv(&self) -> Option<(candle_core::Tensor, candle_core::Tensor)> {
140 match (&self.k_buffer, &self.v_buffer) {
141 (Some(k), Some(v)) if self.current_len > 0 => {
142 let k_view = k.narrow(2, 0, self.current_len).ok()?;
144 let v_view = v.narrow(2, 0, self.current_len).ok()?;
145 Some((k_view, v_view))
146 }
147 _ => None,
148 }
149 }
150
151 #[deprecated(note = "Use get_kv() instead for better performance")]
155 pub fn k(&self) -> Option<Tensor> {
156 self.get_kv().map(|(k, _)| Tensor::from_candle(k))
157 }
158
159 #[deprecated(note = "Use get_kv() instead for better performance")]
161 pub fn v(&self) -> Option<Tensor> {
162 self.get_kv().map(|(_, v)| Tensor::from_candle(v))
163 }
164
165 #[deprecated(note = "Use append() instead for better performance")]
167 pub fn set_k(&mut self, k: Tensor) {
168 if let Ok(candle_k) = k.to_candle() {
169 self.k_buffer = Some(candle_k);
170 if let Some(ref k) = self.k_buffer {
171 self.current_len = k.dims().get(2).copied().unwrap_or(0);
172 }
173 }
174 }
175
176 #[deprecated(note = "Use append() instead for better performance")]
178 pub fn set_v(&mut self, v: Tensor) {
179 if let Ok(candle_v) = v.to_candle() {
180 self.v_buffer = Some(candle_v);
181 }
182 }
183}
184
185impl Default for LayerKVCache {
186 fn default() -> Self {
187 Self::new()
188 }
189}
190
191#[derive(Debug)]
193pub struct KVCache {
194 layers: Vec<LayerKVCache>,
196 seq_len: usize,
198 max_seq_len: usize,
200}
201
202impl KVCache {
203 pub fn new(num_layers: usize) -> Self {
205 Self::with_capacity(num_layers, DEFAULT_MAX_SEQ_LEN)
206 }
207
208 pub fn with_capacity(num_layers: usize, max_seq_len: usize) -> Self {
210 Self {
211 layers: (0..num_layers)
212 .map(|_| LayerKVCache::with_capacity(max_seq_len))
213 .collect(),
214 seq_len: 0,
215 max_seq_len,
216 }
217 }
218
219 pub fn num_layers(&self) -> usize {
221 self.layers.len()
222 }
223
224 pub fn layer_mut(&mut self, layer_idx: usize) -> &mut LayerKVCache {
226 &mut self.layers[layer_idx]
227 }
228
229 pub fn layer(&self, layer_idx: usize) -> &LayerKVCache {
231 &self.layers[layer_idx]
232 }
233
234 pub fn seq_len(&self) -> usize {
236 self.seq_len
237 }
238
239 pub fn set_seq_len(&mut self, seq_len: usize) {
241 self.seq_len = seq_len;
242 }
243
244 pub fn max_seq_len(&self) -> usize {
246 self.max_seq_len
247 }
248
249 pub fn clear(&mut self) {
251 for layer in &mut self.layers {
252 layer.clear();
253 }
254 self.seq_len = 0;
255 }
256
257 pub fn deallocate(&mut self) {
259 for layer in &mut self.layers {
260 layer.deallocate();
261 }
262 self.seq_len = 0;
263 }
264
265 pub fn is_empty(&self) -> bool {
267 self.seq_len == 0
268 }
269}
270
271#[cfg(test)]
272mod tests {
273 use super::*;
274
275 #[test]
276 fn test_kv_cache_creation() {
277 let cache = KVCache::new(32);
278 assert_eq!(cache.num_layers(), 32);
279 assert_eq!(cache.seq_len(), 0);
280 assert!(cache.is_empty());
281 assert_eq!(cache.max_seq_len(), DEFAULT_MAX_SEQ_LEN);
282 }
283
284 #[test]
285 fn test_kv_cache_with_capacity() {
286 let cache = KVCache::with_capacity(16, 4096);
287 assert_eq!(cache.num_layers(), 16);
288 assert_eq!(cache.max_seq_len(), 4096);
289 }
290
291 #[test]
292 fn test_layer_cache_empty() {
293 let layer_cache = LayerKVCache::new();
294 assert!(layer_cache.is_empty());
295 assert_eq!(layer_cache.seq_len(), 0);
296 assert!(layer_cache.get_kv().is_none());
297 }
298
299 #[test]
300 fn test_layer_cache_append() {
301 let mut layer_cache = LayerKVCache::with_capacity(512);
302
303 let k = candle_core::Tensor::zeros(
305 (1, 4, 10, 64),
306 candle_core::DType::F32,
307 &candle_core::Device::Cpu,
308 )
309 .unwrap();
310 let v = candle_core::Tensor::zeros(
311 (1, 4, 10, 64),
312 candle_core::DType::F32,
313 &candle_core::Device::Cpu,
314 )
315 .unwrap();
316
317 layer_cache.append(&k, &v).unwrap();
319 assert_eq!(layer_cache.seq_len(), 10);
320 assert!(!layer_cache.is_empty());
321
322 let (k_view, v_view) = layer_cache.get_kv().unwrap();
324 assert_eq!(k_view.dims(), &[1, 4, 10, 64]);
325 assert_eq!(v_view.dims(), &[1, 4, 10, 64]);
326
327 let k2 = candle_core::Tensor::zeros(
329 (1, 4, 5, 64),
330 candle_core::DType::F32,
331 &candle_core::Device::Cpu,
332 )
333 .unwrap();
334 let v2 = candle_core::Tensor::zeros(
335 (1, 4, 5, 64),
336 candle_core::DType::F32,
337 &candle_core::Device::Cpu,
338 )
339 .unwrap();
340
341 layer_cache.append(&k2, &v2).unwrap();
342 assert_eq!(layer_cache.seq_len(), 15);
343
344 let (k_view, v_view) = layer_cache.get_kv().unwrap();
345 assert_eq!(k_view.dims(), &[1, 4, 15, 64]);
346 assert_eq!(v_view.dims(), &[1, 4, 15, 64]);
347 }
348
349 #[test]
350 fn test_layer_cache_clear() {
351 let mut layer_cache = LayerKVCache::with_capacity(512);
352
353 let k = candle_core::Tensor::zeros(
354 (1, 4, 10, 64),
355 candle_core::DType::F32,
356 &candle_core::Device::Cpu,
357 )
358 .unwrap();
359 let v = candle_core::Tensor::zeros(
360 (1, 4, 10, 64),
361 candle_core::DType::F32,
362 &candle_core::Device::Cpu,
363 )
364 .unwrap();
365
366 layer_cache.append(&k, &v).unwrap();
367 assert_eq!(layer_cache.seq_len(), 10);
368
369 layer_cache.clear();
371 assert_eq!(layer_cache.seq_len(), 0);
372 assert!(layer_cache.is_empty());
373 assert!(layer_cache.k_buffer.is_some()); layer_cache.deallocate();
377 assert!(layer_cache.k_buffer.is_none());
378 }
379
380 #[test]
381 fn test_kv_cache_clear() {
382 let mut cache = KVCache::new(4);
383 cache.set_seq_len(100);
384 assert_eq!(cache.seq_len(), 100);
385
386 cache.clear();
387 assert_eq!(cache.seq_len(), 0);
388 assert!(cache.is_empty());
389 }
390}