Skip to main content

runtime/
kv_cache.rs

1//! KV Cache for efficient autoregressive generation
2//!
3//! This module provides a pre-allocated CPU-focused KV cache that stores computed
4//! Key and Value tensors to avoid recomputation during text generation.
5//!
6//! # How it works
7//!
8//! During autoregressive generation, each transformer layer computes Key and Value
9//! tensors that are reused when generating subsequent tokens. Without caching,
10//! the model must recompute K,V for all previous tokens on each generation step,
11//! leading to O(n^2) complexity.
12//!
13//! With KV caching:
14//! - Prefill: Process entire prompt, cache K,V for each layer
15//! - Decode: For each new token, only compute K,V for that token and append
16//!   to cached values using slice_set (no allocation!)
17//!
18//! This reduces generation to O(n) complexity, providing 50-100x speedup.
19
20use anyhow::Result;
21use crate::tensor_core::Tensor;
22
23/// Default maximum sequence length for pre-allocation
24pub const DEFAULT_MAX_SEQ_LEN: usize = 2048;
25
26/// Per-layer KV cache storing Key and Value tensors with pre-allocated buffers
27#[derive(Debug, Clone)]
28pub struct LayerKVCache {
29    /// Pre-allocated Key buffer: [batch, num_kv_heads, max_seq_len, head_dim]
30    k_buffer: Option<candle_core::Tensor>,
31    /// Pre-allocated Value buffer: [batch, num_kv_heads, max_seq_len, head_dim]
32    v_buffer: Option<candle_core::Tensor>,
33    /// Current filled length (how many positions are used)
34    current_len: usize,
35    /// Maximum sequence length (buffer capacity)
36    max_seq_len: usize,
37}
38
39impl LayerKVCache {
40    /// Create an empty layer cache with default max sequence length
41    pub fn new() -> Self {
42        Self::with_capacity(DEFAULT_MAX_SEQ_LEN)
43    }
44
45    /// Create a layer cache with specified maximum sequence length
46    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    /// Check if this layer has cached values
56    pub fn is_empty(&self) -> bool {
57        self.current_len == 0
58    }
59
60    /// Clear cached values (keeps buffers allocated for reuse)
61    pub fn clear(&mut self) {
62        self.current_len = 0;
63        // Note: We keep buffers allocated for reuse
64    }
65
66    /// Fully deallocate buffers
67    pub fn deallocate(&mut self) {
68        self.k_buffer = None;
69        self.v_buffer = None;
70        self.current_len = 0;
71    }
72
73    /// Get the current cached sequence length
74    pub fn seq_len(&self) -> usize {
75        self.current_len
76    }
77
78    /// Get maximum sequence length (buffer capacity)
79    pub fn max_seq_len(&self) -> usize {
80        self.max_seq_len
81    }
82
83    /// Append new K,V tensors to the cache using slice_set (zero-allocation)
84    ///
85    /// # Arguments
86    /// * `new_k` - New key tensor: [batch, num_kv_heads, new_seq_len, head_dim]
87    /// * `new_v` - New value tensor: [batch, num_kv_heads, new_seq_len, head_dim]
88    ///
89    /// # Returns
90    /// Ok(()) on success, Err if buffer overflow
91    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        // Lazy allocation on first append
95        if self.k_buffer.is_none() {
96            let shape = new_k.dims();
97            // [batch, num_kv_heads, max_seq_len, head_dim]
98            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        // Check for buffer overflow
112        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        // Use slice_set for zero-allocation append
122        // slice_set(&src, dim, start) writes src into buffer at position start on dimension dim
123        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    /// Get view of the filled portion of K,V buffers
137    ///
138    /// Returns (K, V) tensors narrowed to current sequence length
139    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                // narrow(dim, start, len) returns a view, no copy
143                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    // === Legacy compatibility methods ===
152
153    /// Get cached K tensor (legacy interface)
154    #[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    /// Get cached V tensor (legacy interface)
160    #[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    /// Set K tensor (legacy interface - converts to append)
166    #[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    /// Set V tensor (legacy interface)
177    #[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/// Full model KV cache containing all layers
192#[derive(Debug)]
193pub struct KVCache {
194    /// Per-layer caches
195    layers: Vec<LayerKVCache>,
196    /// Current total sequence length in cache
197    seq_len: usize,
198    /// Maximum sequence length for all layers
199    max_seq_len: usize,
200}
201
202impl KVCache {
203    /// Create a new KV cache with default max sequence length
204    pub fn new(num_layers: usize) -> Self {
205        Self::with_capacity(num_layers, DEFAULT_MAX_SEQ_LEN)
206    }
207
208    /// Create a new KV cache with specified max sequence length
209    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    /// Get the number of layers
220    pub fn num_layers(&self) -> usize {
221        self.layers.len()
222    }
223
224    /// Get mutable reference to a specific layer's cache
225    pub fn layer_mut(&mut self, layer_idx: usize) -> &mut LayerKVCache {
226        &mut self.layers[layer_idx]
227    }
228
229    /// Get reference to a specific layer's cache
230    pub fn layer(&self, layer_idx: usize) -> &LayerKVCache {
231        &self.layers[layer_idx]
232    }
233
234    /// Get current cached sequence length
235    pub fn seq_len(&self) -> usize {
236        self.seq_len
237    }
238
239    /// Update the cached sequence length
240    pub fn set_seq_len(&mut self, seq_len: usize) {
241        self.seq_len = seq_len;
242    }
243
244    /// Get maximum sequence length
245    pub fn max_seq_len(&self) -> usize {
246        self.max_seq_len
247    }
248
249    /// Clear all cached values (keeps buffers allocated for reuse)
250    pub fn clear(&mut self) {
251        for layer in &mut self.layers {
252            layer.clear();
253        }
254        self.seq_len = 0;
255    }
256
257    /// Fully deallocate all buffers
258    pub fn deallocate(&mut self) {
259        for layer in &mut self.layers {
260            layer.deallocate();
261        }
262        self.seq_len = 0;
263    }
264
265    /// Check if cache is empty
266    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        // Create test tensors: [batch=1, heads=4, seq=10, head_dim=64]
304        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        // Append first batch
318        layer_cache.append(&k, &v).unwrap();
319        assert_eq!(layer_cache.seq_len(), 10);
320        assert!(!layer_cache.is_empty());
321
322        // Get K,V view
323        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        // Append more tokens
328        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        // Clear resets length but keeps buffers
370        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()); // Buffer still allocated
374
375        // Deallocate removes buffers
376        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}