#!/usr/bin/env python3 """H2O KV-cache eviction for DeepseekV3Attention / Multi-head Latent Attention (MLA). HuggingFace's reference implementation of MLA caches the *expanded* K/V tensors (not the compressed latent), so KV memory scales as: L layers * seq_len * H heads * (qk_dim + v_dim) dims * 2 bytes (FP16) For a model with the canonical DeepseekV3 layout (61 layers, 64 heads, qk_dim=192, v_dim=128), the cache reaches: ~82 GB at 32K context ~328 GB at 128K context This module patches every DeepseekV3Attention layer with H2O-style "heavy-hitter + recency" eviction: - Keep the first `n_sink` tokens unconditionally (attention sinks) - Track per-token accumulated attention mass across all heads/layers - When cache exceeds `budget`, evict the lowest-score tokens (excluding sinks + recent) Usage: from kv_eviction_mla import install_kv_eviction, remove_kv_eviction model = AutoModelForCausalLM.from_pretrained(...) install_kv_eviction(model, budget=4096, n_sink=4, n_recent=256) # run inference normally; eviction happens automatically remove_kv_eviction(model) # restore original forward Args: budget : max KV tokens to keep per layer (excluding sinks + recent) n_sink : first N tokens kept unconditionally (attention sink effect) n_recent : last N tokens always kept (ensures fluent local context) evict_every : evict once every N new tokens (amortises overhead; 1 = every step) Memory model at budget=4096 on the canonical DeepseekV3 layout: 61 layers * 4096 tokens * 40960 bytes per token ~= 10.2 GB (versus 82 GB at 32K full cache, 328 GB at 128K full cache) Compatible architectures (verified against HF reference implementations): - DeepSeek V3 / V3.2 family - Kimi K2 / K2.6 (uses DeepseekV3Attention layer pattern) - Any future MLA-based model that subclasses DeepseekV3Attention For non-MLA architectures (standard MHA / GQA), the same H2O recipe applies but the cache-management code paths differ; this module would need adaptation. References: H2O (Zhang et al., 2023): https://arxiv.org/abs/2306.14048 StreamingLLM (Xiao et al., 2023): https://arxiv.org/abs/2309.17453 SnapKV (Li et al., 2024): https://arxiv.org/abs/2404.14469 License: Apache 2.0 """ from __future__ import annotations import types from typing import Optional, Tuple import torch import torch.nn as nn # ────────────────────────────────────────────────────────────────────────────── # Per-layer eviction state # ────────────────────────────────────────────────────────────────────────────── class _EvictionState: """Mutable state attached to each attention layer.""" def __init__(self, budget: int, n_sink: int, n_recent: int, evict_every: int): self.budget = budget self.n_sink = n_sink self.n_recent = n_recent self.evict_every = evict_every # Accumulated attention mass per cached token; updated each forward call. self.score: torch.Tensor | None = None # (bsz, kv_len) self.steps_since_evict: int = 0 def reset(self): self.score = None self.steps_since_evict = 0 # ────────────────────────────────────────────────────────────────────────────── # Patched forward # ────────────────────────────────────────────────────────────────────────────── def _make_evicting_forward(original_forward, state: _EvictionState): """Return a forward method that wraps original_forward with H2O eviction.""" def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value=None, output_attentions: bool = False, use_cache: bool = False, **kwargs, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: # ── Evict BEFORE this forward if cache is over budget ──────────────── if ( past_key_value is not None and use_cache and state.score is not None ): state.steps_since_evict += 1 if state.steps_since_evict >= state.evict_every: state.steps_since_evict = 0 _maybe_evict(past_key_value, self.layer_idx, state) # ── Run original forward (forces output_attentions=True internally) ── # We need attn_weights to update scores; request them explicitly. # original_forward is the BOUND method, so do not pass self. result = original_forward( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value, output_attentions=True, # always get weights for scoring use_cache=use_cache, **kwargs, ) # Unpack defensively: newer transformers may return (attn_out, attn_weights) # while older returns (attn_out, attn_weights, pkv). The cache lives on the # past_key_value object we passed in (mutated in-place in modern versions), # so pkv being absent is fine. if isinstance(result, tuple): if len(result) == 3: attn_out, attn_weights, pkv = result elif len(result) == 2: attn_out, attn_weights = result pkv = past_key_value else: attn_out, attn_weights, pkv = result[0], None, past_key_value else: attn_out, attn_weights, pkv = result, None, past_key_value # ── Update per-token importance scores ─────────────────────────────── if attn_weights is not None and use_cache: # attn_weights: (bsz, num_heads, q_len, kv_len) # Accumulate mean attention mass across heads and current query tokens new_mass = attn_weights.detach().float().mean(dim=(1, 2)) # (bsz, kv_len) kv_len = new_mass.shape[-1] if state.score is None or state.score.shape[-1] != kv_len: state.score = new_mass else: state.score = state.score + new_mass # Return attn_weights=None if caller didn't request them if not output_attentions: attn_weights = None # Return the same arity as the original_forward returned, so calling # decoder layers (which expect 2-tuple in modern transformers, 3-tuple # in older versions) unpack it correctly. if isinstance(result, tuple): if len(result) == 2: return attn_out, attn_weights elif len(result) == 3: return attn_out, attn_weights, pkv return attn_out, attn_weights return forward def _maybe_evict(past_key_value, layer_idx: int, state: _EvictionState) -> None: """Trim the KV cache for layer_idx if it exceeds budget.""" try: # HuggingFace DynamicCache stores lists indexed by layer k = past_key_value.key_cache[layer_idx] # (bsz, heads, kv_len, head_dim) v = past_key_value.value_cache[layer_idx] except (AttributeError, IndexError): return # cache format not supported — skip silently bsz, heads, kv_len, _ = k.shape keep_total = state.n_sink + state.budget + state.n_recent if kv_len <= keep_total: return # still within budget score = state.score # (bsz, kv_len) or None if score is None: return # Indices never evicted: first n_sink and last n_recent protected = set(range(state.n_sink)) | set(range(kv_len - state.n_recent, kv_len)) evictable = [i for i in range(state.n_sink, kv_len - state.n_recent) if i not in protected] if not evictable: return # Pick top-budget evictable positions by score (mean across batch) evict_scores = score[:, evictable].mean(dim=0) # (n_evictable,) n_keep = min(state.budget, len(evictable)) _, top_idx = evict_scores.topk(n_keep, largest=True, sorted=False) keep_evictable = sorted(top_idx.tolist()) # Build final keep indices: sinks + top heavy-hitters + recents sink_idx = list(range(state.n_sink)) recent_idx = list(range(kv_len - state.n_recent, kv_len)) keep_idx = sorted(set(sink_idx + [evictable[i] for i in keep_evictable] + recent_idx)) keep_t = torch.tensor(keep_idx, device=k.device, dtype=torch.long) past_key_value.key_cache[layer_idx] = k[:, :, keep_t, :] past_key_value.value_cache[layer_idx] = v[:, :, keep_t, :] # Trim score to match new kv_len state.score = score[:, keep_t] # ────────────────────────────────────────────────────────────────────────────── # Public API # ────────────────────────────────────────────────────────────────────────────── _ORIGINAL_FORWARD_ATTR = "_h2o_original_forward" _EVICTION_STATE_ATTR = "_h2o_eviction_state" def install_kv_eviction( model: nn.Module, budget: int = 4096, n_sink: int = 4, n_recent: int = 256, evict_every: int = 1, verbose: bool = True, ) -> int: """Patch all DeepseekV3Attention layers in model with H2O KV eviction. Returns the number of layers patched. """ try: from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING except ImportError: pass # Find attention layers by class name (works for any model using DeepseekV3Attention) patched = 0 for name, module in model.named_modules(): cls_name = type(module).__name__ if "Attention" not in cls_name: continue # Target DeepseekV3Attention and its flash variant if cls_name not in ("DeepseekV3Attention", "DeepseekV3FlashAttention2", "DeepseekV3SdpaAttention"): continue if hasattr(module, _ORIGINAL_FORWARD_ATTR): continue # already patched state = _EvictionState(budget, n_sink, n_recent, evict_every) orig = module.forward setattr(module, _ORIGINAL_FORWARD_ATTR, orig) setattr(module, _EVICTION_STATE_ATTR, state) module.forward = types.MethodType(_make_evicting_forward(orig, state), module) patched += 1 if verbose: if patched: kv_gb = patched * (budget + n_sink + n_recent) * 40960 / 1e9 print( f"[kimi_kv_eviction] Patched {patched} attention layers. " f"Budget: {budget} tokens/layer + {n_sink} sinks + {n_recent} recent. " f"Estimated peak KV RAM: {kv_gb:.1f} GB " f"(was ~{patched * 32768 * 40960 / 1e9:.0f} GB at 32K ctx)." ) else: print("[kimi_kv_eviction] No DeepseekV3Attention layers found — model not patched.") return patched def remove_kv_eviction(model: nn.Module) -> int: """Restore original forward methods. Returns number of layers restored.""" restored = 0 for module in model.modules(): if hasattr(module, _ORIGINAL_FORWARD_ATTR): module.forward = getattr(module, _ORIGINAL_FORWARD_ATTR) delattr(module, _ORIGINAL_FORWARD_ATTR) delattr(module, _EVICTION_STATE_ATTR) restored += 1 return restored def reset_eviction_scores(model: nn.Module) -> None: """Clear accumulated importance scores (call between independent generations).""" for module in model.modules(): if hasattr(module, _EVICTION_STATE_ATTR): getattr(module, _EVICTION_STATE_ATTR).reset() # ────────────────────────────────────────────────────────────────────────────── # Quick smoke-test (runs without Kimi K2 — uses a tiny fake model) # ────────────────────────────────────────────────────────────────────────────── if __name__ == "__main__": print("Smoke test: checking patch logic on mock attention layer...") class _FakeCache: def __init__(self, bsz, heads, seq, kd, vd, device): self.key_cache = [torch.randn(bsz, heads, seq, kd, device=device)] self.value_cache = [torch.randn(bsz, heads, seq, vd, device=device)] state = _EvictionState(budget=8, n_sink=2, n_recent=2, evict_every=1) state.score = torch.rand(1, 20) # 20 tokens in cache cache = _FakeCache(1, 4, 20, 192, 128, "cpu") _maybe_evict(cache, 0, state) kept = cache.key_cache[0].shape[2] assert kept == 2 + 8 + 2, f"Expected 12 kept, got {kept}" print(f" OK — evicted 20→{kept} tokens (2 sinks + 8 heavy-hitters + 2 recent)") print("\nUsage example:") print(" from kv_eviction_mla import install_kv_eviction, reset_eviction_scores") print(" install_kv_eviction(model, budget=4096, n_sink=4, n_recent=512)") print(" # generate ...") print(" reset_eviction_scores(model) # between conversations")