""" LeanLlama: Llama with learned KV-cache compression. Load with: AutoModelForCausalLM.from_pretrained("LeanLlama-8B", trust_remote_code=True) """ from __future__ import annotations 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 convert it) # --------------------------------------------------------------------------- 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) # --------------------------------------------------------------------------- # Compressor module (mirrors training-time KVVectorCompressor architecture) # --------------------------------------------------------------------------- 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 # --------------------------------------------------------------------------- # Tensor reshaping helpers (per_token granularity: all heads as one vector) # --------------------------------------------------------------------------- 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 stub — 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+"] _keep_in_fp32_modules = ["leanllm_v_compressor", "leanllm_k_compressor"] 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: KV-cache value compression --- 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] = {} # Track how many tokens have been compressed per layer to avoid re-compressing 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 # Attach value compressor at the same path used by convert_to_hf_model.py # so that from_pretrained auto-loads the saved weights. 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 # Key compressor (if not values-only) 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 # ------------------------------------------------------------------ # KV cache compression (applied after each forward pass) # ------------------------------------------------------------------ def _compress_values( self, v: torch.Tensor, layer_idx: int, start: int, ) -> torch.Tensor: """Compress values from position `start` onward, return full 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: """Compress keys from position `start` onward, return full 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(self, past_key_values: Any) -> Any: if past_key_values is None or not self._kv_layers: return past_key_values # DynamicCache path (transformers >= 4.36) if isinstance(past_key_values, Cache): for layer_idx in self._kv_layers: layer_cache = past_key_values.layers[layer_idx] v = layer_cache.values # [B, H, T, D] 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 # Legacy tuple path (older transformers versions) if isinstance(past_key_values, tuple): past_list = list(past_key_values) for layer_idx in self._kv_layers: 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 def forward(self, *args: Any, **kwargs: Any) -> CausalLMOutputWithPast: # Reset compression tracking when there's no cache (new sequence) past = kwargs.get("past_key_values", None) if past is None and len(args) < 5: # No KV cache passed — starting fresh for layer_idx in self._kv_layers: self._compressed_up_to[layer_idx] = 0 outputs = super().forward(*args, **kwargs) if hasattr(outputs, "past_key_values"): outputs.past_key_values = self._compress_past(outputs.past_key_values) return outputs