Spaces:
Sleeping
Sleeping
| """ | |
| Key-Value Cache for autoregressive generation. | |
| During generation each new token only needs to attend to its own Q | |
| against the full history of K and V. Rather than recomputing K/V for | |
| the entire sequence every step, we cache them per layer and append | |
| the new token's K/V pair. This reduces generation from O(n^2) to | |
| O(n) in sequence length per step. | |
| This module provides TWO implementations of the same interface, kept | |
| side-by-side for benchmarking and educational comparison: | |
| 1. ``KVCacheNaive`` — the textbook implementation. Each ``append`` call | |
| does ``torch.cat([prev_k, k_new], dim=2)``, which allocates a fresh | |
| tensor and copies the entire history every step. Simple to read, | |
| but appends are O(n_history) so total generation cost is O(n^2) | |
| *just from cache management*, on top of the unavoidable O(n) per | |
| token from attention itself. | |
| 2. ``KVCachePrealloc`` — the optimized version. Per-layer fixed-shape | |
| buffers of capacity ``(batch, n_kv_heads, capacity, head_dim)`` | |
| are allocated lazily on first append. Subsequent appends are | |
| in-place slice writes, so they are O(n_new) regardless of how | |
| much history already exists. Buffers grow geometrically (double | |
| on overflow) when the optional ``max_seq_len`` is not pinned, | |
| giving amortized O(1) appends with no upfront waste. Truncation | |
| and reset become O(1) — they only adjust per-layer length | |
| counters; the underlying buffers are reused across runs. | |
| Both classes expose the same public API: | |
| append(layer_idx, k_new, v_new) -> (k_view, v_view) | |
| get(layer_idx) -> (k, v) | None | |
| update(layer_idx, (k, v)) | |
| truncate(layer_idx, length) | |
| truncate_all(length) | |
| trim_to_last(layer_idx, max_len) | |
| reset() | |
| advance(num_tokens) | |
| seq_len (property) | |
| length(layer_idx) | |
| so the kernel and model code is generic over which cache is in use. | |
| ``KVCache`` is aliased to ``KVCachePrealloc`` for the default import. | |
| """ | |
| import torch | |
| class KVCacheNaive: | |
| """ | |
| Educational baseline cache: stores ``(k, v)`` tuples per layer and | |
| grows them via ``torch.cat`` on every append. | |
| Kept in the codebase as the reference implementation that the | |
| pre-allocated version is benchmarked against — re-allocating and | |
| copying the full history each step makes its per-step cost grow | |
| linearly with the sequence length already cached. | |
| """ | |
| def __init__(self, n_layers, **_unused): | |
| self.n_layers = n_layers | |
| self.cache = [None] * n_layers | |
| self.current_pos = 0 | |
| # --- position tracking --- | |
| def seq_len(self): | |
| return self.current_pos | |
| def advance(self, num_tokens): | |
| self.current_pos += num_tokens | |
| def length(self, layer_idx): | |
| entry = self.cache[layer_idx] | |
| return 0 if entry is None else entry[0].size(2) | |
| # --- core append (O(n_history) — the whole point of the comparison) --- | |
| def append(self, layer_idx, k_new, v_new): | |
| prev = self.cache[layer_idx] | |
| if prev is None: | |
| k, v = k_new, v_new | |
| else: | |
| prev_k, prev_v = prev | |
| k = torch.cat([prev_k, k_new], dim=2) | |
| v = torch.cat([prev_v, v_new], dim=2) | |
| self.cache[layer_idx] = (k, v) | |
| return k, v | |
| # --- back-compat tuple API --- | |
| def get(self, layer_idx): | |
| return self.cache[layer_idx] | |
| def update(self, layer_idx, value): | |
| self.cache[layer_idx] = value | |
| def get_all(self): | |
| return self.cache | |
| # --- truncation / reset --- | |
| def truncate(self, layer_idx, length): | |
| if self.cache[layer_idx] is not None: | |
| k, v = self.cache[layer_idx] | |
| self.cache[layer_idx] = (k[:, :, :length, :], v[:, :, :length, :]) | |
| def truncate_all(self, length): | |
| for i in range(self.n_layers): | |
| self.truncate(i, length) | |
| self.current_pos = length | |
| def trim_to_last(self, layer_idx, max_len): | |
| entry = self.cache[layer_idx] | |
| if entry is None: | |
| return | |
| k, v = entry | |
| if k.size(2) <= max_len: | |
| return | |
| self.cache[layer_idx] = (k[:, :, -max_len:, :], v[:, :, -max_len:, :]) | |
| def reset(self): | |
| for i in range(self.n_layers): | |
| self.cache[i] = None | |
| self.current_pos = 0 | |
| class KVCachePrealloc: | |
| """ | |
| Pre-allocated KV cache: per-layer fixed buffers with in-place writes. | |
| Buffers are allocated lazily on the first ``append`` (the cache | |
| does not need to know dtype/shape/device up front) and grow | |
| geometrically when their capacity is exceeded — appends are | |
| amortized O(1). Truncation and reset only flip per-layer length | |
| counters, so speculative-decoding rollbacks become free and | |
| buffers are reused across generation runs. | |
| """ | |
| def __init__(self, n_layers, max_seq_len=None, initial_capacity=256): | |
| """ | |
| Args: | |
| n_layers: Number of transformer layers. | |
| max_seq_len: Optional fixed maximum capacity. If given, buffers | |
| are allocated once at this capacity and never grow; an | |
| append that would overflow raises ``RuntimeError``. | |
| initial_capacity: Starting capacity when ``max_seq_len`` is None. | |
| Buffers double in capacity on overflow. | |
| """ | |
| self.n_layers = n_layers | |
| self.max_seq_len = max_seq_len | |
| self.initial_capacity = initial_capacity | |
| self.k_buffers = [None] * n_layers | |
| self.v_buffers = [None] * n_layers | |
| self.layer_lens = [0] * n_layers | |
| self.current_pos = 0 | |
| # --- position tracking --- | |
| def seq_len(self): | |
| return self.current_pos | |
| def advance(self, num_tokens): | |
| self.current_pos += num_tokens | |
| def length(self, layer_idx): | |
| return self.layer_lens[layer_idx] | |
| # --- core append (in-place, O(n_new)) --- | |
| def append(self, layer_idx, k_new, v_new): | |
| """ | |
| Append new K/V tensors for ``layer_idx`` and return views over the | |
| full prefix (length = previous + new). | |
| Both ``k_new`` and ``v_new`` have shape | |
| ``(batch, n_kv_heads, n_new_tokens, head_dim)``. | |
| """ | |
| self._ensure_capacity(layer_idx, k_new, v_new) | |
| k_buf = self.k_buffers[layer_idx] | |
| v_buf = self.v_buffers[layer_idx] | |
| cur = self.layer_lens[layer_idx] | |
| n = k_new.size(2) | |
| k_buf[:, :, cur:cur + n, :] = k_new | |
| v_buf[:, :, cur:cur + n, :] = v_new | |
| new_len = cur + n | |
| self.layer_lens[layer_idx] = new_len | |
| return k_buf[:, :, :new_len, :], v_buf[:, :, :new_len, :] | |
| # --- back-compat tuple API --- | |
| def get(self, layer_idx): | |
| n = self.layer_lens[layer_idx] | |
| if n == 0 or self.k_buffers[layer_idx] is None: | |
| return None | |
| k_buf = self.k_buffers[layer_idx] | |
| v_buf = self.v_buffers[layer_idx] | |
| return k_buf[:, :, :n, :], v_buf[:, :, :n, :] | |
| def update(self, layer_idx, value): | |
| """ | |
| Replace the stored K/V for ``layer_idx`` with ``value`` (a ``(k, v)`` | |
| tuple). Copies the data into the pre-allocated buffer; reallocates | |
| the buffer if it does not exist or is too small. | |
| """ | |
| if value is None: | |
| self.layer_lens[layer_idx] = 0 | |
| return | |
| k_new, v_new = value | |
| n = k_new.size(2) | |
| # Reset length so the ensure-capacity check sizes for the full payload | |
| self.layer_lens[layer_idx] = 0 | |
| self._ensure_capacity(layer_idx, k_new, v_new) | |
| k_buf = self.k_buffers[layer_idx] | |
| v_buf = self.v_buffers[layer_idx] | |
| k_buf[:, :, :n, :] = k_new | |
| v_buf[:, :, :n, :] = v_new | |
| self.layer_lens[layer_idx] = n | |
| def get_all(self): | |
| return [self.get(i) for i in range(self.n_layers)] | |
| # --- truncation / reset (O(1) — just adjust length counters) --- | |
| def truncate(self, layer_idx, length): | |
| if self.k_buffers[layer_idx] is None: | |
| return | |
| self.layer_lens[layer_idx] = min(self.layer_lens[layer_idx], length) | |
| def truncate_all(self, length): | |
| for i in range(self.n_layers): | |
| self.truncate(i, length) | |
| self.current_pos = length | |
| def trim_to_last(self, layer_idx, max_len): | |
| """ | |
| Keep only the most recent ``max_len`` tokens for ``layer_idx``. | |
| Used by sliding-window attention: when the cached length exceeds the | |
| window, slide the kept slice to the front of the buffer in-place so | |
| the next ``append`` writes contiguously after it. | |
| """ | |
| cur = self.layer_lens[layer_idx] | |
| if cur <= max_len or self.k_buffers[layer_idx] is None: | |
| return | |
| k_buf = self.k_buffers[layer_idx] | |
| v_buf = self.v_buffers[layer_idx] | |
| # Clone the slice we want to keep before overwriting (avoid aliasing). | |
| keep_k = k_buf[:, :, cur - max_len:cur, :].clone() | |
| keep_v = v_buf[:, :, cur - max_len:cur, :].clone() | |
| k_buf[:, :, :max_len, :] = keep_k | |
| v_buf[:, :, :max_len, :] = keep_v | |
| self.layer_lens[layer_idx] = max_len | |
| def reset(self): | |
| """Clear lengths and global position. Buffers are kept for reuse.""" | |
| for i in range(self.n_layers): | |
| self.layer_lens[i] = 0 | |
| self.current_pos = 0 | |
| # --- internal: lazy buffer allocation and geometric growth --- | |
| def _ensure_capacity(self, layer_idx, k_new, v_new): | |
| n_new = k_new.size(2) | |
| cur = self.layer_lens[layer_idx] | |
| needed = cur + n_new | |
| k_buf = self.k_buffers[layer_idx] | |
| if k_buf is None: | |
| cap = self.max_seq_len if self.max_seq_len is not None else max( | |
| self.initial_capacity, n_new | |
| ) | |
| shape = (k_new.size(0), k_new.size(1), cap, k_new.size(3)) | |
| self.k_buffers[layer_idx] = torch.empty( | |
| shape, dtype=k_new.dtype, device=k_new.device | |
| ) | |
| self.v_buffers[layer_idx] = torch.empty( | |
| shape, dtype=v_new.dtype, device=v_new.device | |
| ) | |
| return | |
| if needed <= k_buf.size(2): | |
| return | |
| if self.max_seq_len is not None: | |
| raise RuntimeError( | |
| f"KVCachePrealloc layer {layer_idx} would exceed " | |
| f"max_seq_len={self.max_seq_len} (needed {needed})." | |
| ) | |
| # Grow geometrically — at least double, at least enough to fit `needed`. | |
| new_cap = max(needed, k_buf.size(2) * 2) | |
| shape = (k_buf.size(0), k_buf.size(1), new_cap, k_buf.size(3)) | |
| new_k = torch.empty(shape, dtype=k_buf.dtype, device=k_buf.device) | |
| new_v = torch.empty(shape, dtype=k_buf.dtype, device=k_buf.device) | |
| new_k[:, :, :cur, :] = k_buf[:, :, :cur, :] | |
| new_v[:, :, :cur, :] = self.v_buffers[layer_idx][:, :, :cur, :] | |
| self.k_buffers[layer_idx] = new_k | |
| self.v_buffers[layer_idx] = new_v | |
| # Default cache used by the model code. Aliased so existing imports | |
| # (`from src.utils.kv_cache import KVCache`) pick up the optimized version. | |
| KVCache = KVCachePrealloc | |