diff --git a/README.md b/README.md index 1b425f04..43328127 100644 --- a/README.md +++ b/README.md @@ -128,6 +128,7 @@ Some presses rely on a different logic: - `DuoAttentionPress` ([source](kvpress/presses/duo_attention_press.py), [paper](https://arxiv.org/abs/2410.10819)): split heads into retrieval heads (no compression) and streaming heads (StreamingLLM approach) - `FinchPress` ([source](kvpress/presses/finch_press.py), [paper](https://direct.mit.edu/tacl/article/doi/10.1162/tacl_a_00716/125280)): similar to SnapKV with a dynamic window size and key value re-rotation - `KVzipPress` ([source](kvpress/presses/kvzip_press.py), [paper](https://arxiv.org/abs/2505.23416)): identify redundant KV pairs through context reconstruction. Achieve near-lossless compression at the cost of multiple forward passes. + - `structure_demotion` (KVzipPress and RestoreKVPress option, off by default): before eviction, scale the scores of repeated punctuation/whitespace tokens (list separators, newlines) by this factor. Reconstruction scores rate such scaffolding highly and at high compression ratios it absorbs a large share of the budget; demoting it (e.g. `0.25`) hands that budget to content tokens while never touching letters or digits. - `KVgradPress` ([source](kvpress/presses/kvgrad_press.py), [paper](https://openreview.net/forum?id=cg1wTCJjjk)): similar to `KVzipPress`, but scores KV pairs by an input × gradient attribution of their effect on the last hidden states, at the cost of one backward pass per chunk. - `RestoreKVPress` ([source](kvpress/presses/restorekv_press.py), [paper](https://arxiv.org/abs/2608.01247)): extends the KV cache with 8 learned restore tokens encoded by a LoRA module before budget-matched KVzip pruning. - `KVComposePress` ([source](kvpress/presses/kvcompose_press.py), [paper](https://arxiv.org/abs/2509.05165)): attention-guided eviction, aligning per-head selections into composite tokens to preserve cache structure. diff --git a/evaluation/evaluate_registry.py b/evaluation/evaluate_registry.py index 9552e9b0..4e86e190 100644 --- a/evaluation/evaluate_registry.py +++ b/evaluation/evaluate_registry.py @@ -102,6 +102,7 @@ "kvgrad": KVgradPress(), "kvzip": KVzipPress(), "kvzip_plus": KVzipPress(kvzip_plus_normalization=True), + "kvzip_plus_sd": KVzipPress(kvzip_plus_normalization=True, structure_demotion=0.25), # + structure demotion "kvzap_linear": DMSPress(press=KVzapPress(model_type="linear")), "kvzap_mlp": DMSPress(press=KVzapPress(model_type="mlp")), "kvzap_mlp_head": KVzapPress(model_type="mlp"), @@ -115,6 +116,7 @@ "random": RandomPress(), "RestoreKV": RestoreKVPress(), "RestoreKV_plus": RestoreKVPress(kvzip_plus_normalization=True), # RestoreKV+ (KVzip+ scoring) + "RestoreKV_plus_sd": RestoreKVPress(kvzip_plus_normalization=True, structure_demotion=0.25), # + structure demotion "snap_think": ComposedPress([SnapKVPress(), ThinKPress()]), "snapkv": SnapKVPress(), "streaming_llm": StreamingLLMPress(), diff --git a/kvpress/presses/kvzip_press.py b/kvpress/presses/kvzip_press.py index 3e6bd4d3..ba2e6b2f 100644 --- a/kvpress/presses/kvzip_press.py +++ b/kvpress/presses/kvzip_press.py @@ -3,14 +3,22 @@ import logging import math +import re from contextlib import contextmanager from dataclasses import dataclass from types import MethodType -from typing import Generator, List +from typing import Callable, Generator, List, cast import torch from torch import nn -from transformers import AutoTokenizer, Gemma3PreTrainedModel, PreTrainedModel, PreTrainedTokenizer, QuantizedCache +from transformers import ( + AutoTokenizer, + Gemma3PreTrainedModel, + PreTrainedModel, + PreTrainedTokenizer, + PreTrainedTokenizerBase, + QuantizedCache, +) from kvpress.presses.base_press import SUPPORTED_MODELS, BasePress from kvpress.utils import extract_keys_and_values, get_query_states @@ -46,6 +54,15 @@ class KVzipPress(BasePress): Whether to enable KVzip+ normalization. chunk_size : int, default=2048 Number of context tokens reconstructed by each replay pass. + structure_demotion : float, default=0.0 + Structure demotion (off when 0). Before eviction, multiply the scores of *structural* tokens -- tokens + that decode to punctuation, whitespace or symbols only, no letters or digits -- that occur at least + ``structure_min_repeats`` times in the context by this factor. Reconstruction-based scores rate the + repeated scaffolding of long contexts (list separators, newlines, quotes) highly, and at high compression + ratios it absorbs a large share of the budget; demoting it hands that budget to content tokens. Letters and + digits are never touched, so rare facts (needles, numbers) are safe by construction. + structure_min_repeats : int, default=4 + Minimum number of occurrences of a structural token id in the context for it to be demoted. """ compression_ratio: float = 0.0 @@ -53,10 +70,14 @@ class KVzipPress(BasePress): n_sink: int = 4 kvzip_plus_normalization: bool = False chunk_size: int = 2048 + structure_demotion: float = 0.0 + structure_min_repeats: int = 4 def __post_init__(self): assert 0 <= self.compression_ratio < 1, "Compression ratio must be between 0 and 1" assert self.chunk_size > 0, "Chunk size must be positive" + assert 0 <= self.structure_demotion < 1, "structure_demotion is a factor in [0, 1); 0 disables it" + assert self.structure_min_repeats >= 1, "structure_min_repeats must be >= 1" logger.warning( "KVzipPress requires multiple forward passes for chunked context reconstruction, " "resulting in a computational overhead of 2–3 times the initial prefilling cost. " @@ -74,6 +95,7 @@ def _reset_internal_parameters(self): self._cache = None self.score_val = None + self._tokenizer: PreTrainedTokenizerBase | None = None self.causal_mask_score = None self.start_idx = 0 self.end_idx = 0 @@ -97,6 +119,7 @@ def __call__(self, model: PreTrainedModel) -> Generator: # Store model reference for later use tokenizer = AutoTokenizer.from_pretrained(model.config.name_or_path) + self._tokenizer = tokenizer # Get suffix_ids directly using tokenizer's chat template (do this once, not in hook) if tokenizer.chat_template is None: @@ -361,12 +384,91 @@ def score_kvzip( keys, values = keys[:, :, : self.context_length], values[:, :, : self.context_length] return keys, values + # A content token is a word, word piece or number: after stripping whitespace it starts with a letter or digit + # and contains only letters, digits, apostrophes and hyphens. Everything else -- punctuation, whitespace, list + # markers, brackets, quotes -- is structure and may be demoted when highly repeated. + _CONTENT_TOKEN = re.compile(r"[^\W_](?:[^\W_]|['\-])*", re.UNICODE) + _ALNUM_CHAR = re.compile(r"[^\W_]", re.UNICODE) + + @staticmethod + def demote_structure_scores( + score_val: torch.Tensor, + context_ids: torch.Tensor, + decode: Callable[[int], str], + factor: float, + min_repeats: int, + start: int = 0, + ) -> int: + """ + In-place structure demotion on ``score_val`` (n_layer, bsz, n_kv_heads, ctx_len). + + A position is demoted when (a) its token is structural (not a word, word piece or number, see + ``_CONTENT_TOKEN``), (b) the same token id occurs at least ``min_repeats`` times in ``context_ids[start:]``, + and (c) it is not a *joiner*: a whitespace-free structural token whose previous token ends with a letter or + digit and whose next token starts with one -- the dashes of a UUID, the dot of ``3.14``, the colon of + ``10:30``, the comma of ``1,000``. Separators (list markers, quotes, brackets, tokens carrying a newline) + have whitespace on at least one side and are demoted. Positions before ``start`` (chat prefix, sinks) and + beyond ``len(context_ids)`` (e.g. RestoreKV's restore tokens) are never touched. + ``decode(token_id) -> str``. Returns the number of demoted positions. Pure function of tensors + decode. + """ + n = min(int(score_val.shape[-1]), int(context_ids.shape[-1])) + if factor <= 0 or n <= start: + return 0 + ids_all = context_ids[:n].detach().to("cpu") + ids = ids_all[start:] + uniq, inverse, counts = torch.unique(ids, return_inverse=True, return_counts=True) + dec = {int(t): decode(int(t)) for t in uniq.tolist()} + structural = torch.tensor( + [KVzipPress._CONTENT_TOKEN.fullmatch(dec[int(t)].strip()) is None for t in uniq.tolist()], dtype=torch.bool + ) + demote = (structural & (counts >= min_repeats))[inverse] + # joiner exclusion: depends on the neighbours, so it is evaluated per position + toks = ids_all.tolist() + alnum = KVzipPress._ALNUM_CHAR + for j in torch.nonzero(demote).flatten().tolist(): + i = j + start + d = dec[toks[i]] + if i == 0 or i + 1 >= n or any(ch.isspace() for ch in d): + continue + prev_d = dec.get(toks[i - 1]) + if prev_d is None: + prev_d = dec[toks[i - 1]] = decode(toks[i - 1]) + next_d = dec[toks[i + 1]] + if alnum.fullmatch(prev_d[-1:]) and alnum.fullmatch(next_d[:1]): + demote[j] = False + idx = torch.nonzero(demote).flatten() + start + if idx.numel() == 0: + return 0 + idx = idx.to(score_val.device) + score_val[..., idx] = score_val[..., idx] * factor + return int(idx.numel()) + def compress_post(self, model: PreTrainedModel): """ Obtain the indices of KV pairs to be evicted. Adopted from adakv_press.compress (fake compression). KVzip does not rely on safeguards. """ if self.compression_ratio > 0: + if self.structure_demotion > 0 and self._context_ids is not None and self._tokenizer is not None: + tokenizer = self._tokenizer + + def decode(token_id: int) -> str: + return cast(str, tokenizer.decode([token_id])) + + n_demoted = self.demote_structure_scores( + self.score_val, + self._context_ids[0], + decode, + self.structure_demotion, + self.structure_min_repeats, + start=max(self.prefix_length, self.n_sink), + ) + logger.debug( + "Structure demotion: %d of %d positions scaled by %.2f", + n_demoted, + self.score_val.shape[-1], + self.structure_demotion, + ) # Attention sinks are never evicted: they are given the highest score self.score_val[..., : self.n_sink] = self.score_val.amax() + 1.0 n_layer, bsz, num_key_value_heads, ctx_len = self.score_val.shape diff --git a/tests/presses/test_kvzip_structure_demotion.py b/tests/presses/test_kvzip_structure_demotion.py new file mode 100644 index 00000000..3fa6ff0c --- /dev/null +++ b/tests/presses/test_kvzip_structure_demotion.py @@ -0,0 +1,147 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 Satyajeeth Suresh Kannan. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch + +from kvpress.presses.kvzip_press import KVzipPress +from kvpress.presses.restorekv_press import RestoreKVPress + +# a toy vocabulary: 0-3 sinks/prefix tokens, 10-19 words, 20-29 digits, 30-39 structure +VOCAB = {**{i: f"" for i in range(4)}, **{10 + i: f" word{i}" for i in range(10)}} +VOCAB.update({20 + i: str(i) for i in range(10)}) +VOCAB.update( + { + 30: ".", + 31: "\n", + 32: " (", + 33: ",", + 34: " ", + 35: " -", + 36: "'s", + 37: " é", + 38: ")", + 39: "!", + 40: "-", + 41: "-ce", + 42: " word", + 43: ".\n", + } +) + + +def decode(token_id: int) -> str: + return VOCAB[token_id] + + +def _scores(n_positions: int, n_layer: int = 2, n_heads: int = 3) -> torch.Tensor: + return torch.ones(n_layer, 1, n_heads, n_positions) + + +def test_repeated_structure_is_demoted_and_content_is_not(): + # prefix(2) | "1. word0\n" x 5 with distinct words, plus a needle number that repeats as a token + ids = [0, 1] + for i in range(5): + ids += [20 + i, 30, 10 + i, 31] # digit, ".", word, "\n" + ids += [27, 27, 27, 27, 27] # the digit 7 five times: content, never demoted + context = torch.tensor(ids) + score = _scores(len(ids)) + n = KVzipPress.demote_structure_scores(score, context, decode, factor=0.25, min_repeats=4, start=2) + demoted = torch.nonzero(score[0, 0, 0] < 1).flatten().tolist() + expected = [i for i, t in enumerate(ids) if i >= 2 and t in (30, 31)] # "." and "\n" occur 5x each + assert demoted == expected + assert n == len(expected) + assert torch.allclose(score[..., expected], torch.full((2, 1, 3, len(expected)), 0.25)) + # digits, words and the prefix are untouched, even the digit repeated 5 times + kept = [i for i in range(len(ids)) if i not in expected] + assert torch.all(score[..., kept] == 1) + + +def test_min_repeats_and_start_are_respected(): + ids = [30, 30, 30, 1, 30, 30, 30, 33, 33, 33] # "." x3 in the region, "," x3 + context = torch.tensor(ids) + score = _scores(len(ids)) + # region starts at 3: "." occurs 3 times there, "," 3 times -> with min_repeats=4 nothing is demoted + assert KVzipPress.demote_structure_scores(score, context, decode, 0.5, min_repeats=4, start=3) == 0 + assert torch.all(score == 1) + # with min_repeats=3 both are demoted, but never the positions before start + assert KVzipPress.demote_structure_scores(score, context, decode, 0.5, min_repeats=3, start=3) == 6 + assert torch.all(score[..., :3] == 1) + assert torch.all(score[..., 4:] == 0.5) + + +def test_positions_beyond_the_context_are_never_touched(): + # score_val longer than context_ids (e.g. RestoreKV appended restore tokens): tail untouched + ids = [31] * 6 + context = torch.tensor(ids) + score = _scores(len(ids) + 8) + assert KVzipPress.demote_structure_scores(score, context, decode, 0.25, min_repeats=4, start=0) == 6 + assert torch.all(score[..., 6:] == 1) + assert torch.all(score[..., :6] == 0.25) + + +@pytest.mark.parametrize( + "token_id,is_content", + [ + (10, True), + (20, True), + (37, True), + (30, False), + (31, False), + (32, False), + (34, False), + (35, False), + (36, False), + (40, False), + (41, False), + ], +) +def test_token_classification(token_id, is_content): + assert (KVzipPress._CONTENT_TOKEN.fullmatch(decode(token_id).strip()) is not None) is is_content + + +def test_joiners_inside_uuids_are_not_demoted_but_separators_are(): + # "8f-4a-9e" style ids: a dash glued between alphanumerics is a joiner and is never demoted, however often it recurs + uuid = [20, 21, 40, 22, 23, 40, 24, 25, 40, 26, 27] # digits and dashes glued together + ids = uuid + [31] + uuid + [31] + uuid + [31] + uuid + [31] # 4 ids on 4 lines: dash x12, newline x4 + context = torch.tensor(ids) + score = _scores(len(ids)) + n = KVzipPress.demote_structure_scores(score, context, decode, factor=0.25, min_repeats=4, start=0) + dashes = [i for i, t in enumerate(ids) if t == 40] + newlines = [i for i, t in enumerate(ids) if t == 31] + assert torch.all(score[..., dashes] == 1), "UUID dashes must not be demoted" + assert torch.all(score[..., newlines] == 0.25) and n == len(newlines) + # a dash+letters piece ("-ce") glued inside an id is a joiner too + ids = [20, 41, 21] * 5 + score = _scores(len(ids)) + assert ( + KVzipPress.demote_structure_scores(score, torch.tensor(ids), decode, factor=0.25, min_repeats=4, start=0) == 0 + ) + # list markers "1." followed by " word" keep their demotion: the next token starts with whitespace + lst = [] + for i in range(5): + lst += [20 + i, 30, 42, 31] + score = _scores(len(lst)) + KVzipPress.demote_structure_scores(score, torch.tensor(lst), decode, factor=0.25, min_repeats=4, start=0) + dots = [i for i, t in enumerate(lst) if t == 30] + assert torch.all(score[..., dots] == 0.25) + # a token carrying a newline (".\n") between two words is a separator, not a joiner + para = [10, 43, 11] * 5 + score = _scores(len(para)) + KVzipPress.demote_structure_scores(score, torch.tensor(para), decode, factor=0.25, min_repeats=4, start=0) + assert torch.all(score[..., [i for i, t in enumerate(para) if t == 43]] == 0.25) + # a dash at the very first or last position has no two neighbours and is demoted like any structure token + edge = [40, 20, 40, 21, 40, 22, 40] + score = _scores(len(edge)) + KVzipPress.demote_structure_scores(score, torch.tensor(edge), decode, factor=0.25, min_repeats=4, start=0) + assert torch.all(score[..., [0, 6]] == 0.25) and torch.all(score[..., [2, 4]] == 1) + + +def test_factor_zero_disables_and_option_is_inherited_by_restorekv(): + score = _scores(6) + assert KVzipPress.demote_structure_scores(score, torch.tensor([31] * 6), decode, 0.0, 4) == 0 + assert torch.all(score == 1) + press = RestoreKVPress(kvzip_plus_normalization=True, structure_demotion=0.25) + assert press.structure_demotion == 0.25 and press.structure_min_repeats == 4 + with pytest.raises(AssertionError): + KVzipPress(structure_demotion=1.0)