Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 2 additions & 0 deletions evaluation/evaluate_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Expand All @@ -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(),
Expand Down
106 changes: 104 additions & 2 deletions kvpress/presses/kvzip_press.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -46,17 +54,30 @@ 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
layerwise: bool = False
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. "
Expand All @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down
147 changes: 147 additions & 0 deletions tests/presses/test_kvzip_structure_demotion.py
Original file line number Diff line number Diff line change
@@ -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"<s{i}>" 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)