Download model/kv_cache.py from imshadow0/pycraft-1: direct link, hf CLI and curl.
- Browser
- Download file 12.1 kB
-
https://huggingface.co/imshadow0/pycraft-1/resolve/main/model/kv_cache.py
- Command line
-
hf download hf://imshadow0/pycraft-1/model/kv_cache.py
-
curl -L -o kv_cache.py https://huggingface.co/imshadow0/pycraft-1/resolve/main/model/kv_cache.py
12.1 kB
| # model/kv_cache.py | |
| # | |
| # KV-cache for PyCraft-1 incremental decoding. | |
| # | |
| # Without a cache, generating token N re-runs the full forward pass over all | |
| # N-1 previous tokens — quadratic work for a linear amount of output. | |
| # This module stores each layer's keys and values so every decode step only | |
| # computes the single newest position. | |
| # | |
| # Two design choices matter and are easy to get wrong: | |
| # | |
| # 1. Storage is PRE-expansion — (B, n_kv_heads, ...) not (B, n_heads, ...). | |
| # With GQA 8Q/2KV that is 4x less memory, which is the entire point of | |
| # grouped-query attention. _repeat_kv re-expands on read. | |
| # | |
| # 2. Memory is PREALLOCATED, never torch.cat'd. Concatenating per step costs | |
| # O(N) copy per step and O(N^2) overall — roughly 4.3 GB of pointless | |
| # memcpy across one full-context generation. | |
| # | |
| # The cache holds POST-RoPE keys. RoPE is relative (q_m . k_n depends only on | |
| # m - n), so a key rotated once at its true absolute position stays valid | |
| # forever. Never re-apply RoPE to cached keys. | |
| import torch | |
| # ------------------------------------------------------------------ # | |
| # KV cache | |
| # ------------------------------------------------------------------ # | |
| class KVCache: | |
| """ | |
| Pre-allocated per-layer key/value store for incremental decoding. | |
| Per layer: k[i], v[i] : (batch, n_kv_heads, max_len, head_dim) | |
| Live region: [:, :, :seq_len, :] | |
| Typical use: | |
| cache = KVCache.from_model(model, batch_size=1) | |
| logits, _ = model(prompt_ids, past_key_values=cache) # prefill | |
| logits, _ = model(next_id, past_key_values=cache) # decode | |
| """ | |
| def __init__( | |
| self, | |
| n_layers: int, | |
| batch_size: int, | |
| n_kv_heads: int, | |
| head_dim: int, | |
| max_len: int, | |
| device, | |
| dtype: torch.dtype = torch.float32, | |
| ): | |
| self.n_layers = n_layers | |
| self.batch_size = batch_size | |
| self.n_kv_heads = n_kv_heads | |
| self.head_dim = head_dim | |
| self.max_len = max_len | |
| self.device = device | |
| self.dtype = dtype | |
| self.seq_len = 0 | |
| shape = (batch_size, n_kv_heads, max_len, head_dim) | |
| self.k = [torch.zeros(shape, device=device, dtype=dtype) | |
| for _ in range(n_layers)] | |
| self.v = [torch.zeros(shape, device=device, dtype=dtype) | |
| for _ in range(n_layers)] | |
| # -------------------------------------------------------------- # | |
| # Called once per layer, per forward pass | |
| # -------------------------------------------------------------- # | |
| def update( | |
| self, | |
| layer_idx: int, | |
| k_new: torch.Tensor, # (batch, n_kv_heads, T, head_dim) | |
| v_new: torch.Tensor, # (batch, n_kv_heads, T, head_dim) | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """ | |
| Write this layer's new K/V into the cache and return views covering | |
| everything written so far: (batch, n_kv_heads, seq_len + T, head_dim). | |
| Deliberately does NOT advance seq_len. The model advances it once, | |
| after all layers have written. Advancing here would make layer i see | |
| the offset that only layer i+1 should see — a silent corruption that | |
| produces plausible-looking garbage. | |
| """ | |
| batch, _, T, _ = k_new.shape | |
| start = self.seq_len | |
| if start + T > self.max_len: | |
| raise ValueError( | |
| f"KVCache overflow: {start} + {T} > max_len={self.max_len}. " | |
| f"Allocate a larger cache or stop generating." | |
| ) | |
| if batch != self.batch_size: | |
| raise ValueError( | |
| f"batch {batch} does not match cache batch {self.batch_size}" | |
| ) | |
| if k_new.dtype != self.dtype: | |
| # Silent casting here would quietly degrade precision every step. | |
| raise TypeError( | |
| f"cache dtype {self.dtype} != incoming dtype {k_new.dtype}" | |
| ) | |
| self.k[layer_idx][:, :, start:start + T] = k_new | |
| self.v[layer_idx][:, :, start:start + T] = v_new | |
| # Narrowed views — the zero-padded tail is never visible, so reset() | |
| # needs no memset. | |
| return ( | |
| self.k[layer_idx][:, :, :start + T], | |
| self.v[layer_idx][:, :, :start + T], | |
| ) | |
| def advance(self, n: int): | |
| """Advance the write head. Called once per forward, by the model.""" | |
| self.seq_len += n | |
| def reset(self): | |
| """Reuse this cache for a new sequence. No zeroing required.""" | |
| self.seq_len = 0 | |
| def __len__(self) -> int: | |
| return self.seq_len | |
| def __repr__(self) -> str: | |
| return ( | |
| f"KVCache(layers={self.n_layers}, batch={self.batch_size}, " | |
| f"kv_heads={self.n_kv_heads}, head_dim={self.head_dim}, " | |
| f"seq_len={self.seq_len}/{self.max_len}, dtype={self.dtype})" | |
| ) | |
| def memory_bytes(self) -> int: | |
| """Total allocated cache size in bytes (K and V, all layers).""" | |
| per = (self.batch_size * self.n_kv_heads | |
| * self.max_len * self.head_dim) | |
| return 2 * self.n_layers * per * torch.empty( | |
| (), dtype=self.dtype).element_size() | |
| # -------------------------------------------------------------- # | |
| def from_model( | |
| cls, | |
| model, | |
| batch_size: int = 1, | |
| max_len: int | None = None, | |
| device=None, | |
| dtype: torch.dtype | None = None, | |
| ) -> "KVCache": | |
| """Build a cache matching a model's config, device, and dtype.""" | |
| cfg = model.config | |
| # Read device/dtype from the embedding, not a Linear: dynamic | |
| # quantization replaces nn.Linear with a module whose .weight is a | |
| # method. Activations stay fp32 under dynamic quant either way. | |
| ref = model.token_embedding.weight | |
| return cls( | |
| n_layers=cfg.n_layers, | |
| batch_size=batch_size, | |
| n_kv_heads=cfg.n_kv_heads, | |
| head_dim=cfg.head_dim, | |
| max_len=max_len or cfg.max_seq_len, | |
| device=device if device is not None else ref.device, | |
| dtype=dtype if dtype is not None else ref.dtype, | |
| ) | |
| # ------------------------------------------------------------------ # | |
| # Causal mask construction | |
| # ------------------------------------------------------------------ # | |
| def build_attn_mask( | |
| q_len: int, | |
| offset: int, | |
| device, | |
| padding_mask: torch.Tensor | None = None, | |
| ) -> torch.Tensor | None: | |
| """ | |
| Build a bottom-right-aligned causal mask for cached attention. | |
| Args: | |
| padding_mask: optional (batch, offset + q_len) bool marking real | |
| tokens True and padding False. Required for batched generation, | |
| where shorter prompts are left-padded to a common length. | |
| With no padding mask, returns None for the two cases SDPA handles on its | |
| own: | |
| offset == 0 caller passes is_causal=True (square, so PyTorch's | |
| upper-left alignment is already correct) | |
| q_len == 1 the single newest query may attend to every cached | |
| key, so causality holds structurally — no mask needed | |
| Otherwise returns bool "may attend", (1, 1, q_len, offset + q_len) or | |
| (batch, 1, q_len, offset + q_len) when padding is involved: | |
| mask[i, j] = (j <= offset + i) and not padding[j] | |
| WHY THIS EXISTS: PyTorch builds is_causal as torch.ones(L, S).tril() — | |
| UPPER-LEFT aligned. With L=1, S=N that mask contains exactly one True | |
| (column 0), so a cached decode step with is_causal=True would attend only | |
| to the first prompt token, at every layer. No error, no NaN — just output | |
| that ignores the prompt and collapses into repetition. Hence the explicit | |
| dispatch rather than a blanket is_causal=True. | |
| """ | |
| if padding_mask is None and (offset == 0 or q_len == 1): | |
| return None | |
| kv_len = offset + q_len | |
| q_pos = torch.arange(offset, kv_len, device=device).unsqueeze(1) # (T, 1) | |
| k_pos = torch.arange(0, kv_len, device=device).unsqueeze(0) # (1, S) | |
| mask = (k_pos <= q_pos)[None, None, :, :] # (1,1,T,S) | |
| if padding_mask is not None: | |
| if padding_mask.shape[-1] != kv_len: | |
| raise ValueError( | |
| f"padding_mask covers {padding_mask.shape[-1]} positions but " | |
| f"the key length is {kv_len} (offset {offset} + q_len {q_len})" | |
| ) | |
| # (B,1,1,S) broadcast against the causal (1,1,T,S) -> (B,1,T,S) | |
| mask = mask & padding_mask[:, None, None, :].to(torch.bool) | |
| # A fully-masked query row makes softmax(all -inf) = NaN, which then | |
| # spreads through the residual stream. Left-padded rows hit this: a | |
| # pad token's own row can have every key masked. Letting each query | |
| # attend to at least its own position costs nothing (those outputs are | |
| # discarded) and keeps the tensor finite. | |
| qi = torch.arange(q_len, device=device) | |
| mask = mask.clone() | |
| mask[:, :, qi, offset + qi] = True | |
| return mask | |
| # ------------------------------------------------------------------ # | |
| # Quick self-test | |
| # ------------------------------------------------------------------ # | |
| if __name__ == "__main__": | |
| print("Testing build_attn_mask...") | |
| # Fast paths return None | |
| assert build_attn_mask(4, 0, "cpu") is None, "offset=0 should return None" | |
| assert build_attn_mask(1, 7, "cpu") is None, "q_len=1 should return None" | |
| # Explicit mask matches a brute-force construction | |
| for q_len, offset in [(4, 7), (3, 1), (2, 2)]: | |
| got = build_attn_mask(q_len, offset, "cpu") | |
| kv_len = offset + q_len | |
| assert got.shape == (1, 1, q_len, kv_len), f"bad shape {got.shape}" | |
| for i in range(q_len): | |
| for j in range(kv_len): | |
| expected = (j <= offset + i) | |
| assert bool(got[0, 0, i, j]) == expected, ( | |
| f"mask[{i}][{j}] wrong for q_len={q_len}, offset={offset}" | |
| ) | |
| print(" build_attn_mask: OK") | |
| print("\nTesting KVCache...") | |
| cache = KVCache(n_layers=2, batch_size=1, n_kv_heads=2, | |
| head_dim=64, max_len=16, device="cpu") | |
| print(f" {cache}") | |
| print(f" allocated: {cache.memory_bytes / 1024:.1f} KiB") | |
| # Prefill 4 positions across both layers | |
| k = torch.randn(1, 2, 4, 64) | |
| v = torch.randn(1, 2, 4, 64) | |
| for layer in range(2): | |
| kk, vv = cache.update(layer, k, v) | |
| assert kk.shape == (1, 2, 4, 64), f"bad view shape {kk.shape}" | |
| assert cache.seq_len == 0, "update() must not advance seq_len" | |
| cache.advance(4) | |
| assert len(cache) == 4 | |
| # One decode step | |
| k1 = torch.randn(1, 2, 1, 64) | |
| v1 = torch.randn(1, 2, 1, 64) | |
| kk, vv = cache.update(0, k1, v1) | |
| assert kk.shape == (1, 2, 5, 64), f"bad decode view {kk.shape}" | |
| assert torch.equal(kk[:, :, :4], k), "prefill K was corrupted" | |
| assert torch.equal(kk[:, :, 4:], k1), "decode K not written" | |
| cache.advance(1) | |
| print(" writes and views: OK") | |
| # Guards | |
| for bad, exc, label in [ | |
| (lambda: cache.update(0, torch.randn(1, 2, 99, 64), | |
| torch.randn(1, 2, 99, 64)), ValueError, "overflow"), | |
| (lambda: cache.update(0, torch.randn(1, 2, 1, 64).half(), | |
| torch.randn(1, 2, 1, 64).half()), TypeError, "dtype"), | |
| (lambda: cache.update(0, torch.randn(3, 2, 1, 64), | |
| torch.randn(3, 2, 1, 64)), ValueError, "batch"), | |
| ]: | |
| try: | |
| bad() | |
| raise AssertionError(f"{label} guard did not fire") | |
| except exc: | |
| pass | |
| print(" guards: OK") | |
| cache.reset() | |
| assert len(cache) == 0 | |
| print(" reset: OK") | |
| print("\nAll kv_cache tests PASSED.") | |