diff --git a/kvpress/presses/adakv_press.py b/kvpress/presses/adakv_press.py index 7918a6473..fbeea6623 100644 --- a/kvpress/presses/adakv_press.py +++ b/kvpress/presses/adakv_press.py @@ -61,7 +61,7 @@ def compress(self, module, hidden_states, keys, values, attentions, kwargs): bsz, num_key_value_heads, k_len = scores.shape # Make sure to keep at least alpha * (1 - compression_ratio) KV pairs per head - n_kept = int(k_len * (1 - self.compression_ratio)) # ScorerPress definition + n_kept = max(1, int(k_len * (1 - self.compression_ratio))) # ScorerPress definition n_safe = int(n_kept * self.alpha_safeguard) top_indices = torch.topk(scores, n_safe, dim=-1).indices scores.scatter_(-1, top_indices, torch.finfo(scores.dtype).max) diff --git a/kvpress/presses/block_press.py b/kvpress/presses/block_press.py index 22e2da099..1856962e9 100644 --- a/kvpress/presses/block_press.py +++ b/kvpress/presses/block_press.py @@ -63,7 +63,7 @@ def compress( bsz, num_key_value_heads, k_len, head_dim = keys.shape block_size = self.block_size if self.block_size < k_len else k_len - n_kept = int(k_len * (1 - self.compression_ratio)) + n_kept = max(1, int(k_len * (1 - self.compression_ratio))) kept_indices = torch.arange(n_kept, device=keys.device).expand(bsz, num_key_value_heads, -1) diff --git a/kvpress/presses/criticalkv_press.py b/kvpress/presses/criticalkv_press.py index aee4157d1..05c6357b8 100644 --- a/kvpress/presses/criticalkv_press.py +++ b/kvpress/presses/criticalkv_press.py @@ -146,7 +146,7 @@ def compress(self, module, hidden_states, keys, values, attentions, kwargs): bsz, num_key_value_heads, k_len = scores.shape # Make sure to keep at least alpha * (1 - compression_ratio) KV pairs per head - n_kept = int(k_len * (1 - self.compression_ratio)) # ScorerPress definition + n_kept = max(1, int(k_len * (1 - self.compression_ratio))) # ScorerPress definition n_safe = int(n_kept * self.alpha_safeguard) top_indices = torch.topk(scores, n_safe, dim=-1).indices scores.scatter_(-1, top_indices, torch.finfo(scores.dtype).max) diff --git a/kvpress/presses/finch_press.py b/kvpress/presses/finch_press.py index 66879b673..80e4f6f55 100644 --- a/kvpress/presses/finch_press.py +++ b/kvpress/presses/finch_press.py @@ -97,7 +97,7 @@ def compress(self, module, hidden_states, keys, values, attentions, kwargs): # Compute indices to keep (optionally by chunks) k_len = keys.shape[2] # Use actual sequence length from keys instead of hidden_states if self.chunk_length is None: - n_kept = int(k_len * (1 - self.compression_ratio)) + n_kept = max(1, int(k_len * (1 - self.compression_ratio))) indices = scores.topk(n_kept, dim=-1).indices else: assert self.chunk_length > self.window_size / (1 - self.compression_ratio) diff --git a/kvpress/presses/key_rerotation_press.py b/kvpress/presses/key_rerotation_press.py index d79db95a8..b277dbdc3 100644 --- a/kvpress/presses/key_rerotation_press.py +++ b/kvpress/presses/key_rerotation_press.py @@ -143,7 +143,7 @@ def compress( # Get indices of KV pairs with the lowest scores q_len = keys.shape[2] - n_kept = int(q_len * (1 - self.press.compression_ratio)) + n_kept = max(1, int(q_len * (1 - self.press.compression_ratio))) indices = scores.topk(n_kept, dim=-1).indices indices = torch.sort(indices, dim=2).values keys = self.rerotate_keys(module, indices, keys) diff --git a/kvpress/presses/merging_press.py b/kvpress/presses/merging_press.py index ed07ad674..b65bf6b7c 100644 --- a/kvpress/presses/merging_press.py +++ b/kvpress/presses/merging_press.py @@ -83,7 +83,7 @@ def compress( # Get indices of KV pairs with the lowest scores k_len = keys.shape[2] - n_kept = int(k_len * (1 - self.press.compression_ratio)) + n_kept = max(1, int(k_len * (1 - self.press.compression_ratio))) indices = scores.topk(n_kept, dim=-1).indices # Merge evicted tokens into the survivors before pruning diff --git a/kvpress/presses/scorer_press.py b/kvpress/presses/scorer_press.py index f7f33d169..b4e964690 100644 --- a/kvpress/presses/scorer_press.py +++ b/kvpress/presses/scorer_press.py @@ -91,7 +91,7 @@ def compress( # Get indices of KV pairs with the lowest scores k_len = keys.shape[2] - n_kept = int(k_len * (1 - self.compression_ratio)) + n_kept = max(1, int(k_len * (1 - self.compression_ratio))) indices = scores.topk(n_kept, dim=-1).indices indices = indices.unsqueeze(-1).expand(-1, -1, -1, module.head_dim) diff --git a/tests/presses/test_scorer_press.py b/tests/presses/test_scorer_press.py new file mode 100644 index 000000000..d11ffcf55 --- /dev/null +++ b/tests/presses/test_scorer_press.py @@ -0,0 +1,117 @@ +# SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import types + +import pytest +import torch + +from kvpress.presses.adakv_press import AdaKVPress +from kvpress.presses.block_press import BlockPress +from kvpress.presses.criticalkv_press import CriticalAdaKVPress +from kvpress.presses.finch_press import FinchPress +from kvpress.presses.key_rerotation_press import KeyRerotationPress +from kvpress.presses.knorm_press import KnormPress +from kvpress.presses.merging_press import MergingPress + +# Short contexts where ``int(k_len * (1 - compression_ratio))`` floors to 0. +# Without a floor guard the cache is emptied to (bsz, heads, 0, head_dim) and +# the next decode silently attends zero keys (NaN / divergent generation). The +# ``max(1, int(...))`` guard keeps at least one token at every n_kept site. +PRESS_NAMES = [ + "KnormPress", + "BlockPress", + "MergingPress", + "KeyRerotationPress", + "FinchPress", + "AdaKVPress", + "CriticalAdaKVPress", +] + + +def _make_press(name, ratio): + child = KnormPress(compression_ratio=ratio) + if name == "KnormPress": + return child + if name == "BlockPress": + return BlockPress(press=child) + if name == "MergingPress": + return MergingPress(press=child) + if name == "KeyRerotationPress": + return KeyRerotationPress(press=child) + if name == "FinchPress": + press = FinchPress(compression_ratio=ratio, rerotate_keys=False) + press.window_size = 1 + return press + if name == "AdaKVPress": + return AdaKVPress(press=child) + if name == "CriticalAdaKVPress": + return CriticalAdaKVPress(press=child) + raise ValueError(name) + + +def _make_module(name, head_dim, num_key_value_heads): + """Minimal fake attention module exposing only what each press reads.""" + module = types.SimpleNamespace(head_dim=head_dim) + if name == "KeyRerotationPress": + # rerotate_keys reads module.rotary_emb.inv_freq + module.rotary_emb = types.SimpleNamespace(inv_freq=torch.randn(head_dim // 2)) + if name == "FinchPress": + module.config = types.SimpleNamespace(num_attention_heads=num_key_value_heads) + if name in ("AdaKVPress", "CriticalAdaKVPress"): + num_attention_heads = num_key_value_heads # num_key_value_groups == 1 + hidden_size = num_attention_heads * head_dim + module.config = types.SimpleNamespace( + _attn_implementation="sdpa", + num_attention_heads=num_attention_heads, + head_dim=head_dim, + hidden_size=hidden_size, + ) + module.num_key_value_groups = 1 + module.o_proj = types.SimpleNamespace( + weight=torch.randn(hidden_size, num_attention_heads * head_dim) + ) + return module + + +@pytest.mark.parametrize("k_len, ratio", [(1, 0.5), (2, 0.6)]) +@pytest.mark.parametrize("press_name", PRESS_NAMES) +def test_compress_never_empties_cache_on_short_context(press_name, k_len, ratio): + """Short contexts must never collapse the cache to zero tokens. + + Reproduces the n_kept floor-to-zero bug: ``int(k_len * (1 - ratio))`` is 0 + for the parametrized cases, so the cache is emptied (shape[2] == 0) or every + token is flagged for pruning. The ``max(1, ...)`` guard keeps at least one + token, so the compressed cache stays non-empty and finite. + """ + if press_name == "FinchPress" and k_len < 2: + # FinchPress windowing needs a window strictly shorter than the sequence. + pytest.skip("FinchPress requires k_len > window_size") + + bsz, num_key_value_heads, head_dim = 1, 2, 8 + hidden_dim = num_key_value_heads * head_dim + keys = torch.randn(bsz, num_key_value_heads, k_len, head_dim) + values = torch.randn(bsz, num_key_value_heads, k_len, head_dim) + hidden_states = torch.randn(bsz, k_len, hidden_dim) + + module = _make_module(press_name, head_dim, num_key_value_heads) + press = _make_press(press_name, ratio) + + if press_name == "FinchPress": + num_heads = num_key_value_heads # num_key_value_groups == 1 + attentions = torch.randn(bsz, num_heads, press.window_size, k_len) + else: + attentions = None + kwargs = {} + + out_keys, _ = press.compress(module, hidden_states, keys, values, attentions, kwargs) + + # The cache must never be emptied to zero tokens. + assert out_keys.shape[2] >= 1 + assert not torch.isnan(out_keys).any() + + # AdaKVPress / CriticalAdaKVPress return keys unchanged but flag pruned + # tokens via module.masked_key_indices; with the bug every token is pruned. + masked = getattr(module, "masked_key_indices", None) + if masked is not None: + assert len(masked[2]) < num_key_value_heads * k_len