""" LeanLlama v2: Llama with adaptive KV-cache compression. Extends v1 (fixed-rate compression) with attention-based importance scoring: - High-attention tokens → light compression (wide bottleneck) - Low-attention tokens → aggressive compression (narrow bottleneck) - Negligible tokens → evicted from cache entirely Load with: AutoModelForCausalLM.from_pretrained("LeanLlama-v2", trust_remote_code=True) """ from __future__ import annotations from pathlib import Path from typing import Any import torch from torch import nn from transformers import LlamaConfig, LlamaForCausalLM from transformers.cache_utils import Cache from transformers.modeling_outputs import CausalLMOutputWithPast # --------------------------------------------------------------------------- # Quantization-safe linear (not nn.Linear, so bitsandbytes won't quantize) # --------------------------------------------------------------------------- class _SafeLinear(nn.Module): """Drop-in replacement for nn.Linear that bitsandbytes will not quantize.""" def __init__(self, in_features: int, out_features: int, bias: bool = True) -> None: super().__init__() self.in_features = in_features self.out_features = out_features self.weight = nn.Parameter(torch.empty(out_features, in_features)) self.bias = nn.Parameter(torch.empty(out_features)) if bias else None nn.init.kaiming_uniform_(self.weight, a=5**0.5) if self.bias is not None: nn.init.zeros_(self.bias) def forward(self, x: torch.Tensor) -> torch.Tensor: return nn.functional.linear(x, self.weight, self.bias) # --------------------------------------------------------------------------- # Fixed-rate compressor (v1 — unchanged from LeanLlama-8B-INT4) # --------------------------------------------------------------------------- class KVVectorCompressorModule(nn.Module): def __init__( self, input_dim: int, bottleneck_dim: int, decoder_depth: int = 2, decoder_hidden_dim: int = 128, residual_decoder: bool = False, fixed_alpha: float | None = 1.0, ) -> None: super().__init__() self.input_dim = input_dim self.bottleneck_dim = bottleneck_dim self.residual_decoder = residual_decoder self.encoder = _SafeLinear(input_dim, bottleneck_dim) if decoder_depth == 1: self.decoder = nn.Sequential(_SafeLinear(bottleneck_dim, input_dim)) else: self.decoder = nn.Sequential( _SafeLinear(bottleneck_dim, decoder_hidden_dim), nn.GELU(), _SafeLinear(decoder_hidden_dim, input_dim), ) if fixed_alpha is None: self.alpha = nn.Parameter(torch.tensor(1.0, dtype=torch.float32)) self.register_buffer("_alpha_const", torch.tensor(0.0, dtype=torch.float32)) else: self.register_parameter("alpha", None) self.register_buffer("_alpha_const", torch.tensor(float(fixed_alpha), dtype=torch.float32)) def _alpha(self) -> torch.Tensor: return self.alpha if self.alpha is not None else self._alpha_const def encode(self, x: torch.Tensor) -> torch.Tensor: return self.encoder(x) def decode(self, z: torch.Tensor, x_orig: torch.Tensor | None = None) -> torch.Tensor: out = self.decoder(z) if self.residual_decoder and x_orig is not None: return x_orig + self._alpha() * out return out # --------------------------------------------------------------------------- # Adaptive compressor with soft dimension masking (v2) # --------------------------------------------------------------------------- class AdaptiveKVVectorCompressorModule(nn.Module): """Variable-rate KV compressor using soft dimension masking. A single encoder maps to max_bottleneck_dim dimensions. A learned dim_priority parameter determines which dimensions are most important. At encoding time, importance scores control how many dimensions are active via a soft sigmoid mask: importance=1.0 → all max_bottleneck_dim dims active (light compression) importance=0.0 → only min_bottleneck_dim dims active (aggressive compression) """ def __init__( self, input_dim: int, max_bottleneck_dim: int = 128, min_bottleneck_dim: int = 16, decoder_depth: int = 2, decoder_hidden_dim: int = 256, fixed_alpha: float | None = 1.0, mask_temperature: float = 5.0, ) -> None: super().__init__() self.input_dim = input_dim self.max_bottleneck_dim = max_bottleneck_dim self.min_bottleneck_dim = min_bottleneck_dim self.mask_temperature = mask_temperature self.encoder = _SafeLinear(input_dim, max_bottleneck_dim) # Learned dimension priority: which latent dims matter most self.dim_priority = nn.Parameter( torch.linspace(1.0, -1.0, max_bottleneck_dim) ) if decoder_depth == 1: self.decoder = nn.Sequential(_SafeLinear(max_bottleneck_dim, input_dim)) else: self.decoder = nn.Sequential( _SafeLinear(max_bottleneck_dim, decoder_hidden_dim), nn.GELU(), _SafeLinear(decoder_hidden_dim, input_dim), ) if fixed_alpha is None: self.alpha = nn.Parameter(torch.tensor(1.0, dtype=torch.float32)) self.register_buffer("_alpha_const", torch.tensor(0.0, dtype=torch.float32)) else: self.register_parameter("alpha", None) self.register_buffer("_alpha_const", torch.tensor(float(fixed_alpha), dtype=torch.float32)) def _alpha_val(self) -> torch.Tensor: return self.alpha if self.alpha is not None else self._alpha_const def compute_mask(self, importance: torch.Tensor) -> torch.Tensor: """Compute soft dimension mask based on importance scores. Args: importance: [N] tensor with values in [0, 1]. Returns: [N, max_bottleneck_dim] soft mask with values in [0, 1]. """ n_active = ( self.min_bottleneck_dim + importance.unsqueeze(-1) * (self.max_bottleneck_dim - self.min_bottleneck_dim) ) _, priority_order = self.dim_priority.sort(descending=True) rank = torch.empty_like(priority_order) rank[priority_order] = torch.arange( len(priority_order), device=priority_order.device ) rank_float = rank.float().unsqueeze(0) mask = torch.sigmoid( self.mask_temperature * (n_active - rank_float - 0.5) ) return mask def effective_dims(self, importance: torch.Tensor) -> torch.Tensor: """Approximate number of active dimensions per token.""" return self.compute_mask(importance).sum(dim=-1) def encode( self, x: torch.Tensor, importance: torch.Tensor | None = None ) -> torch.Tensor: z = self.encoder(x) if importance is not None: z = z * self.compute_mask(importance) return z def decode(self, z: torch.Tensor, x_orig: torch.Tensor | None = None) -> torch.Tensor: out = self.decoder(z) if x_orig is not None: return x_orig + self._alpha_val() * out return out # --------------------------------------------------------------------------- # Attention-based importance tracking # --------------------------------------------------------------------------- class AdaptiveCompressionState: """Tracks cumulative attention per layer for importance scoring. During autoregressive generation, each new token's attention weights [B, H, 1, T] tell us how much the current query attends to each past key. We accumulate these over time — tokens that are frequently attended to are important. """ def __init__(self) -> None: self.cumulative_attention: dict[int, torch.Tensor] = {} def update(self, layer_idx: int, attn_weights: torch.Tensor) -> None: """Update with [B, H, 1, T] attention from current query → past keys.""" new_attn = attn_weights.squeeze(2) # [B, H, T] if layer_idx not in self.cumulative_attention: self.cumulative_attention[layer_idx] = new_attn.detach() else: prev = self.cumulative_attention[layer_idx] t_new = new_attn.shape[-1] t_old = prev.shape[-1] if t_new > t_old: pad = torch.zeros( prev.shape[0], prev.shape[1], t_new - t_old, device=prev.device, dtype=prev.dtype, ) prev = torch.cat([prev, pad], dim=-1) self.cumulative_attention[layer_idx] = prev + new_attn.detach() def get_importance(self, layer_idx: int) -> torch.Tensor: """Get normalized importance scores [B, T] in [0, 1].""" cum = self.cumulative_attention[layer_idx] # [B, H, T] mean_attn = cum.mean(dim=1) # [B, T] mn = mean_attn.min(dim=-1, keepdim=True).values mx = mean_attn.max(dim=-1, keepdim=True).values return (mean_attn - mn) / (mx - mn + 1e-8) def get_eviction_mask( self, layer_idx: int, threshold: float, sink_size: int = 4, recent_size: int = 64, ) -> torch.Tensor: """Boolean mask of tokens to KEEP (True=keep, False=evict).""" importance = self.get_importance(layer_idx) t = importance.shape[-1] keep = importance >= threshold if sink_size > 0: keep[:, :sink_size] = True if recent_size > 0: keep[:, max(0, t - recent_size):] = True return keep def reset(self) -> None: self.cumulative_attention.clear() # --------------------------------------------------------------------------- # Tensor reshaping helpers # --------------------------------------------------------------------------- def _kv_to_vec(x: torch.Tensor, granularity: str) -> torch.Tensor: """[B, H, T, D] -> [B*T, H*D] (per_token) or [B*T*H, D] (per_head).""" if granularity == "per_head": return x.permute(0, 2, 1, 3).reshape(-1, x.shape[-1]) return x.permute(0, 2, 1, 3).reshape(-1, x.shape[1] * x.shape[3]) def _vec_to_kv(vec: torch.Tensor, ref: torch.Tensor) -> torch.Tensor: """Inverse of _kv_to_vec: restore to [B, H, T, D].""" b, h, t, d = ref.shape return vec.reshape(b, t, h, d).permute(0, 2, 1, 3).contiguous() # --------------------------------------------------------------------------- # Compressed embedding (Phase 1 — activates when config flag is set) # --------------------------------------------------------------------------- class CompressedTokenEmbedding(nn.Module): def __init__(self, vocab_size: int, compressed_dim: int, d_model: int) -> None: super().__init__() self.latent = nn.Embedding(vocab_size, compressed_dim) self.decoder = nn.Sequential( nn.Linear(compressed_dim, min(256, d_model)), nn.GELU(), nn.Linear(min(256, d_model), d_model), ) def forward(self, input_ids: torch.Tensor) -> torch.Tensor: return self.decoder(self.latent(input_ids)) # --------------------------------------------------------------------------- # Main model class # --------------------------------------------------------------------------- class LeanLlamaForCausalLM(LlamaForCausalLM): config_class = LlamaConfig _keys_to_ignore_on_load_unexpected = [r"_leanllm_meta_l\d+"] # _SafeLinear already prevents bitsandbytes from quantizing compressor modules @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs): """Load LeanLlama: resolves base model, loads compressor weights automatically.""" model_dir = Path(pretrained_model_name_or_path) # If the directory has a config with a base model reference, load base weights from there config_path = model_dir / "config.json" if model_dir.is_dir() else None base_model = None if config_path and config_path.exists(): import json with config_path.open() as f: cfg = json.load(f) base_model = cfg.get("leanllm_base_model") if base_model: # Load config from our directory (has compression settings) from transformers import AutoConfig config = AutoConfig.from_pretrained(str(model_dir), trust_remote_code=True) kwargs["config"] = config # Load weights from the base model on HF Hub model = super().from_pretrained(base_model, *args, **kwargs) else: model = super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs) # Auto-load compressor weights if present comp_path = model_dir / "leanllm_compressors.pt" if model_dir.is_dir() else None if comp_path and comp_path.exists(): state = torch.load(comp_path, map_location="cpu", weights_only=False) missing, unexpected = model.load_state_dict(state, strict=False) real_missing = [k for k in missing if "leanllm" in k] if real_missing: print(f"Warning: missing LeanLLM keys: {real_missing}") print(f"[LeanLlama] Loaded {len(state)} compressor params for {len(model._kv_layers)} layers") # Auto-load reconstructed embeddings if present embed_path = model_dir / "leanllm_reconstructed_embeddings.pt" if model_dir.is_dir() else None if embed_path and embed_path.exists(): embed_data = torch.load(embed_path, map_location="cpu", weights_only=False) embed_layer = model.get_input_embeddings() with torch.no_grad(): embed_layer.weight.copy_( embed_data["weight"].to(embed_layer.weight.dtype) ) print("[LeanLlama] Loaded reconstructed embeddings") return model def __init__(self, config: LlamaConfig) -> None: super().__init__(config) # --- Phase 1: Embedding compression (inactive unless config says so) --- emb_cfg = getattr(config, "leanllm_embedding_compression", None) if emb_cfg and emb_cfg.get("enabled"): cdim = int(emb_cfg["compressed_dim"]) self.model.embed_tokens = CompressedTokenEmbedding( vocab_size=config.vocab_size, compressed_dim=cdim, d_model=config.hidden_size, ) # --- Phase 2: Fixed-rate KV-cache compression (v1) --- self._kv_layers: list[int] = [ int(x) for x in getattr(config, "leanllm_compressed_kv_layers", []) ] self._values_only: bool = bool(getattr(config, "leanllm_values_only", True)) kv_input_dims: dict[str, int] = dict(getattr(config, "leanllm_kv_input_dims", {})) kv_value_dims: dict[str, int] = dict(getattr(config, "leanllm_kv_value_dims", {})) kv_key_dims: dict[str, int] = dict(getattr(config, "leanllm_kv_key_dims", {})) kv_decoder_depth: dict[str, int] = dict(getattr(config, "leanllm_kv_decoder_depth", {})) kv_decoder_hidden: dict[str, int] = dict(getattr(config, "leanllm_kv_decoder_hidden_dim", {})) kv_residual: dict[str, bool] = dict(getattr(config, "leanllm_kv_residual_decoder", {})) kv_fixed_alpha: dict[str, float | None] = dict(getattr(config, "leanllm_kv_fixed_alpha", {})) kv_granularity: dict[str, str] = dict(getattr(config, "leanllm_kv_granularity", {})) self._kv_granularity_map: dict[int, str] = {} self._compressed_up_to: dict[int, int] = {} for layer_idx in self._kv_layers: key = str(layer_idx) in_dim = int(kv_input_dims[key]) d_depth = int(kv_decoder_depth.get(key, 2)) d_hidden = int(kv_decoder_hidden.get(key, 128)) r_dec = bool(kv_residual.get(key, False)) f_alpha = kv_fixed_alpha.get(key, 1.0) if f_alpha is not None: f_alpha = float(f_alpha) self._kv_granularity_map[layer_idx] = kv_granularity.get(key, "per_token") self._compressed_up_to[layer_idx] = 0 v_dim = int(kv_value_dims[key]) v_mod = KVVectorCompressorModule( input_dim=in_dim, bottleneck_dim=v_dim, decoder_depth=d_depth, decoder_hidden_dim=d_hidden, residual_decoder=r_dec, fixed_alpha=f_alpha, ) self.model.layers[layer_idx].self_attn.leanllm_v_compressor = v_mod if not self._values_only and key in kv_key_dims: k_dim = int(kv_key_dims[key]) k_mod = KVVectorCompressorModule( input_dim=in_dim, bottleneck_dim=k_dim, decoder_depth=d_depth, decoder_hidden_dim=d_hidden, residual_decoder=r_dec, fixed_alpha=f_alpha, ) self.model.layers[layer_idx].self_attn.leanllm_k_compressor = k_mod # --- Phase 2b: Adaptive KV-cache compression (v2) --- self._adaptive_enabled: bool = bool( getattr(config, "leanllm_adaptive_compression", False) ) self._adaptive_layers: list[int] = [ int(x) for x in getattr(config, "leanllm_adaptive_kv_layers", []) ] self._attn_state: AdaptiveCompressionState | None = None self._adaptive_granularity_map: dict[int, str] = {} self._adaptive_compressed_up_to: dict[int, int] = {} # Eviction config eviction_cfg = getattr(config, "leanllm_eviction", None) or {} self._eviction_enabled: bool = bool(eviction_cfg.get("enabled", False)) self._eviction_threshold: float = float(eviction_cfg.get("threshold", 0.05)) self._eviction_sink_size: int = int(eviction_cfg.get("sink_size", 4)) self._eviction_recent_size: int = int(eviction_cfg.get("recent_size", 128)) self._eviction_interval: int = int(eviction_cfg.get("interval", 64)) self._generation_step: int = 0 adaptive_cfg = getattr(config, "leanllm_adaptive_kv_config", {}) or {} for layer_idx in self._adaptive_layers: key = str(layer_idx) layer_cfg = adaptive_cfg.get(key, {}) in_dim = int(layer_cfg.get("input_dim", kv_input_dims.get(key, config.num_key_value_heads * config.head_dim))) max_bn = int(layer_cfg.get("max_bottleneck_dim", 128)) min_bn = int(layer_cfg.get("min_bottleneck_dim", 16)) d_depth = int(layer_cfg.get("decoder_depth", 2)) d_hidden = int(layer_cfg.get("decoder_hidden_dim", 256)) f_alpha = layer_cfg.get("fixed_alpha", 1.0) if f_alpha is not None: f_alpha = float(f_alpha) mask_temp = float(layer_cfg.get("mask_temperature", 5.0)) gran = str(layer_cfg.get("granularity", "per_token")) self._adaptive_granularity_map[layer_idx] = gran self._adaptive_compressed_up_to[layer_idx] = 0 v_mod = AdaptiveKVVectorCompressorModule( input_dim=in_dim, max_bottleneck_dim=max_bn, min_bottleneck_dim=min_bn, decoder_depth=d_depth, decoder_hidden_dim=d_hidden, fixed_alpha=f_alpha, mask_temperature=mask_temp, ) self.model.layers[layer_idx].self_attn.leanllm_adaptive_v_compressor = v_mod if self._adaptive_enabled and self._adaptive_layers: self._attn_state = AdaptiveCompressionState() # ------------------------------------------------------------------ # Fixed-rate compression (v1) # ------------------------------------------------------------------ def _compress_values(self, v: torch.Tensor, layer_idx: int, start: int) -> torch.Tensor: if start >= v.shape[2]: return v v_new = v[:, :, start:, :] gran = self._kv_granularity_map[layer_idx] v_dtype = v.dtype compressor = self.model.layers[layer_idx].self_attn.leanllm_v_compressor comp_dtype = compressor.encoder.weight.dtype v_vec = _kv_to_vec(v_new, gran).to(comp_dtype) v_rec = compressor.decode(compressor.encode(v_vec), x_orig=None) v_compressed = _vec_to_kv(v_rec, v_new).to(dtype=v_dtype) if start == 0: return v_compressed return torch.cat([v[:, :, :start, :], v_compressed], dim=2) def _compress_keys(self, k: torch.Tensor, layer_idx: int, start: int) -> torch.Tensor: if start >= k.shape[2]: return k k_new = k[:, :, start:, :] gran = self._kv_granularity_map[layer_idx] k_dtype = k.dtype k_comp = self.model.layers[layer_idx].self_attn.leanllm_k_compressor comp_dtype = k_comp.encoder.weight.dtype k_vec = _kv_to_vec(k_new, gran).to(comp_dtype) k_rec = k_comp.decode(k_comp.encode(k_vec), x_orig=None) k_compressed = _vec_to_kv(k_rec, k_new).to(dtype=k_dtype) if start == 0: return k_compressed return torch.cat([k[:, :, :start, :], k_compressed], dim=2) def _compress_past_fixed(self, past_key_values: Any) -> Any: """Apply fixed-rate v1 compression to KV cache.""" if past_key_values is None or not self._kv_layers: return past_key_values if isinstance(past_key_values, Cache): for layer_idx in self._kv_layers: if layer_idx in self._adaptive_layers: continue # handled by adaptive path layer_cache = past_key_values.layers[layer_idx] v = layer_cache.values total_tokens = v.shape[2] start = self._compressed_up_to[layer_idx] layer_cache.values = self._compress_values(v, layer_idx, start) if not self._values_only and hasattr( self.model.layers[layer_idx].self_attn, "leanllm_k_compressor" ): layer_cache.keys = self._compress_keys( layer_cache.keys, layer_idx, start ) self._compressed_up_to[layer_idx] = total_tokens return past_key_values if isinstance(past_key_values, tuple): past_list = list(past_key_values) for layer_idx in self._kv_layers: if layer_idx in self._adaptive_layers: continue k, v = past_list[layer_idx] total_tokens = v.shape[2] start = self._compressed_up_to[layer_idx] v_new = self._compress_values(v, layer_idx, start) if not self._values_only and hasattr( self.model.layers[layer_idx].self_attn, "leanllm_k_compressor" ): k_new = self._compress_keys(k, layer_idx, start) else: k_new = k past_list[layer_idx] = (k_new, v_new) self._compressed_up_to[layer_idx] = total_tokens return tuple(past_list) return past_key_values # ------------------------------------------------------------------ # Adaptive compression (v2) # ------------------------------------------------------------------ def _compress_values_adaptive( self, v: torch.Tensor, layer_idx: int, importance: torch.Tensor | None, start: int, ) -> torch.Tensor: """Compress values with importance-based soft masking.""" if start >= v.shape[2]: return v v_new = v[:, :, start:, :] gran = self._adaptive_granularity_map[layer_idx] v_dtype = v.dtype compressor = self.model.layers[layer_idx].self_attn.leanllm_adaptive_v_compressor comp_dtype = compressor.encoder.weight.dtype v_vec = _kv_to_vec(v_new, gran).to(comp_dtype) # Build per-vector importance from per-token importance imp = None if importance is not None: b, h, t_new, d = v_new.shape # importance is [B, T_total] — slice to match new tokens imp_slice = importance[:, start:] # [B, t_new] if gran == "per_head": imp = imp_slice.unsqueeze(2).expand(b, t_new, h).reshape(-1) else: imp = imp_slice.reshape(-1) imp = imp.to(comp_dtype) z = compressor.encode(v_vec, importance=imp) v_rec = compressor.decode(z, x_orig=None) v_compressed = _vec_to_kv(v_rec, v_new).to(dtype=v_dtype) if start == 0: return v_compressed return torch.cat([v[:, :, :start, :], v_compressed], dim=2) def _get_token_importance(self, layer_idx: int) -> torch.Tensor | None: """Get importance scores for all tokens at a layer. For the newest token (which has zero cumulative attention), uses the median importance of existing tokens as a proxy. """ if self._attn_state is None: return None if layer_idx not in self._attn_state.cumulative_attention: return None imp_all = self._attn_state.get_importance(layer_idx) # [B, T] if imp_all.shape[-1] <= 1: return torch.full_like(imp_all, 0.5) # The last token always has importance=0 (it's new, no future queries yet). # Replace it with the median of existing tokens as a proxy. imp_existing = imp_all[:, :-1] median_imp = imp_existing.median(dim=-1, keepdim=True).values imp_fixed = torch.cat([imp_existing, median_imp], dim=-1) return imp_fixed def _compress_past_adaptive( self, past_key_values: Any, attentions: tuple | None, ) -> Any: """Apply adaptive importance-based compression to KV cache.""" if past_key_values is None or not self._adaptive_layers: return past_key_values if self._attn_state is None: return past_key_values # Update cumulative attention state from this forward pass if attentions is not None: for layer_idx in self._adaptive_layers: if layer_idx < len(attentions) and attentions[layer_idx] is not None: self._attn_state.update(layer_idx, attentions[layer_idx]) if isinstance(past_key_values, Cache): for layer_idx in self._adaptive_layers: layer_cache = past_key_values.layers[layer_idx] v = layer_cache.values total_tokens = v.shape[2] start = self._adaptive_compressed_up_to[layer_idx] importance = self._get_token_importance(layer_idx) layer_cache.values = self._compress_values_adaptive( v, layer_idx, importance, start ) self._adaptive_compressed_up_to[layer_idx] = total_tokens # Token eviction if self._eviction_enabled and self._generation_step > 0: if self._generation_step % self._eviction_interval == 0: past_key_values = self._evict_tokens(past_key_values) return past_key_values if isinstance(past_key_values, tuple): past_list = list(past_key_values) for layer_idx in self._adaptive_layers: k, v = past_list[layer_idx] total_tokens = v.shape[2] start = self._adaptive_compressed_up_to[layer_idx] importance = self._get_token_importance(layer_idx) v_new = self._compress_values_adaptive( v, layer_idx, importance, start ) past_list[layer_idx] = (k, v_new) self._adaptive_compressed_up_to[layer_idx] = total_tokens return tuple(past_list) return past_key_values # ------------------------------------------------------------------ # Token eviction # ------------------------------------------------------------------ def _evict_tokens(self, past_key_values: Any) -> Any: """Remove low-importance tokens from the KV cache.""" if not isinstance(past_key_values, Cache): return past_key_values if self._attn_state is None: return past_key_values # Find tokens to keep across all adaptive layers (intersection) keep_mask = None for layer_idx in self._adaptive_layers: if layer_idx not in self._attn_state.cumulative_attention: continue layer_mask = self._attn_state.get_eviction_mask( layer_idx, threshold=self._eviction_threshold, sink_size=self._eviction_sink_size, recent_size=self._eviction_recent_size, ) if keep_mask is None: keep_mask = layer_mask else: keep_mask = keep_mask & layer_mask # conservative: keep if ANY layer needs it if keep_mask is None: return past_key_values # Only evict if we'd actually remove tokens n_keep = int(keep_mask[0].sum().item()) n_total = keep_mask.shape[-1] if n_keep >= n_total: return past_key_values # Apply eviction to ALL layers (both fixed and adaptive) keep_indices = keep_mask[0].nonzero(as_tuple=True)[0] # [n_keep] n_layers = len(past_key_values.layers) for layer_idx in range(n_layers): layer_cache = past_key_values.layers[layer_idx] k = layer_cache.keys # [B, H, T, D] v = layer_cache.values if k.shape[2] != n_total: continue # skip if size mismatch (shouldn't happen) layer_cache.keys = k[:, :, keep_indices, :] layer_cache.values = v[:, :, keep_indices, :] # Reset compression tracking for layer_idx in self._kv_layers: self._compressed_up_to[layer_idx] = n_keep for layer_idx in self._adaptive_layers: self._adaptive_compressed_up_to[layer_idx] = n_keep # Reset attention state (scores are invalidated by eviction) self._attn_state.reset() return past_key_values # ------------------------------------------------------------------ # Forward # ------------------------------------------------------------------ def forward(self, *args: Any, **kwargs: Any) -> CausalLMOutputWithPast: # Reset compression tracking when starting a new sequence past = kwargs.get("past_key_values", None) if past is None and len(args) < 5: for layer_idx in self._kv_layers: self._compressed_up_to[layer_idx] = 0 for layer_idx in self._adaptive_layers: self._adaptive_compressed_up_to[layer_idx] = 0 if self._attn_state is not None: self._attn_state.reset() self._generation_step = 0 # For adaptive compression, request attention weights during generation # (single-token steps only — prefill is too expensive with eager attention) need_attentions = ( self._adaptive_enabled and self._adaptive_layers and past is not None # generation step, not prefill ) if need_attentions: kwargs.setdefault("output_attentions", True) outputs = super().forward(*args, **kwargs) if hasattr(outputs, "past_key_values"): # Apply fixed-rate compression (v1 layers) outputs.past_key_values = self._compress_past_fixed( outputs.past_key_values ) # Apply adaptive compression (v2 layers) if self._adaptive_enabled: attentions = getattr(outputs, "attentions", None) outputs.past_key_values = self._compress_past_adaptive( outputs.past_key_values, attentions ) self._generation_step += 1 return outputs