Instructions to use transformers-community/sink_cache with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use transformers-community/sink_cache with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("transformers-community/sink_cache", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Fix SinkCache RoPE coordinates for absolute generation
Browse files## Problem
In Transformers 4.52.1, the standard generation loop advances Qwen2 RoPE
positions absolutely. The upstream SinkCache truncates its KV window but also
rerotates every surviving key into a bounded local slot. Its stored key phase
therefore no longer matches the new query's absolute phase.
## Change
- The custom `generate()` wrapper selects absolute-position mode. A shift
retains each surviving key's original RoPE phase while dropping old K/V.
- Direct callers can select `position_mode="local"` when both `position_ids`
and `cache_position` intentionally use bounded slots. That path computes
rerotation in FP32, normalizes the derived twiddle, then writes BF16 keys.
- Appending exactly to capacity no longer evicts a token early.
- README documents the position contract and the tested Transformers version.
## Verification
- Qwen2.5-1.5B-Instruct, BF16, eager attention, Transformers 4.52.1, RTX 4080;
two cache arms interleaved against the same frozen model and WikiText-2
teacher-forced token stream. Prefill logits matched bitwise.
- 2,600 tokens, four sinks, 256 recent positions, absolute positions:
2,339 evictions; upstream 65,492 rerotations and 27.840 post-fill PPL;
corrected cache zero rerotations and 19.227 PPL. Every reported block
favored the correction.
- On the identical model/data/tokens/prefill with bounded local positions,
correctly rerotating the keys gave 19.268 PPL upstream and 19.236 with
FP32 normalized rerotation. The corrected absolute and correctly local
contracts therefore agree closely.
- At a 1,024-recent-token window and 20,000 tokens with bounded local
positions, upstream and FP32 normalized rerotation were nearly tied:
10.122 versus 10.136 post-fill PPL. The numerical change improves an
isolated retained key's fidelity but is not a demonstrated next-token
quality gain under that local contract.
- Eleven regression tests pass under Transformers 4.52.1, including native
Qwen2 generation, exact capacity fill, and repeated BF16 key shifts.
- A paired 20,000-token, 1,024-recent-token absolute-position follow-up is
queued on the same GPU; I will add its result to this PR when it completes.
## Scope
The paired model-forward test reproduces the relevant Qwen2 generation
position/cache contract while keeping the input tokens identical between
arms. It is not a throughput benchmark. This historical custom cache is not
compatible with the changed Cache API in Transformers 5.x; the README now
states the tested version.
- README.md +19 -3
- custom_generate/generate.py +38 -9
- tests/test_sink_cache_precision.py +243 -0
|
@@ -11,7 +11,22 @@ This is done by always keeping the first few tokens ("sink tokens") in the KV ca
|
|
| 11 |
amount of attention to them. As it discards past non-sink tokens, the model will lose the ability to generate tokens
|
| 12 |
that depend on the context that was discarded. It's also a solution to contain the memory footprint of the KV cache.
|
| 13 |
|
| 14 |
-
This implementation
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
|
| 16 |

|
| 17 |
|
|
@@ -27,7 +42,8 @@ This implementation matches the `SinkCache` class present in `transformers<4.53.
|
|
| 27 |
|
| 28 |
|
| 29 |
## Additional Arguments
|
| 30 |
-
- `window_length` (`int`, *optional*, defaults to 256): The
|
|
|
|
| 31 |
- `num_sink_tokens` (`int`, *optional*, defaults to 4): The number of sink tokens. See the original paper for more information.
|
| 32 |
|
| 33 |
|
|
@@ -41,7 +57,7 @@ in `generate.py`, in this repository.
|
|
| 41 |
We can use the custom generation method in this repository like the the base `generate` from `transformers`:
|
| 42 |
|
| 43 |
```py
|
| 44 |
-
#
|
| 45 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 46 |
|
| 47 |
# Preparing model, tokenizer, and model inputs
|
|
|
|
| 11 |
amount of attention to them. As it discards past non-sink tokens, the model will lose the ability to generate tokens
|
| 12 |
that depend on the context that was discarded. It's also a solution to contain the memory footprint of the KV cache.
|
| 13 |
|
| 14 |
+
This implementation is based on the `SinkCache` class present in `transformers<4.53.0`.
|
| 15 |
+
The behavior below was verified with `transformers==4.52.1`. The `Cache` API in
|
| 16 |
+
Transformers 5.x is incompatible with this class; the example is not currently
|
| 17 |
+
supported on that version.
|
| 18 |
+
|
| 19 |
+
### RoPE positions when the cache shifts
|
| 20 |
+
|
| 21 |
+
The standard `model.generate()` path advances RoPE positions absolutely. When
|
| 22 |
+
the window discards old entries, surviving keys keep their original rotation;
|
| 23 |
+
only the stored K/V entries move. Direct callers that intentionally reuse
|
| 24 |
+
bounded local `position_ids` and `cache_position` instead need the surviving
|
| 25 |
+
keys rerotated to their new local positions. Keep those two position inputs in
|
| 26 |
+
the same coordinate system. The custom `generate()` wrapper sets absolute mode
|
| 27 |
+
explicitly. Direct `SinkCache` users can pass `position_mode="absolute"` or
|
| 28 |
+
`position_mode="local"`; the default `"auto"` infers the mode at first eviction
|
| 29 |
+
from `cache_position` and assumes positions started at zero.
|
| 30 |
|
| 31 |

|
| 32 |
|
|
|
|
| 42 |
|
| 43 |
|
| 44 |
## Additional Arguments
|
| 45 |
+
- `window_length` (`int`, *optional*, defaults to 256): The total KV cache capacity,
|
| 46 |
+
including `num_sink_tokens`. For four sinks and 1,024 recent tokens, set this to 1,028.
|
| 47 |
- `num_sink_tokens` (`int`, *optional*, defaults to 4): The number of sink tokens. See the original paper for more information.
|
| 48 |
|
| 49 |
|
|
|
|
| 57 |
We can use the custom generation method in this repository like the the base `generate` from `transformers`:
|
| 58 |
|
| 59 |
```py
|
| 60 |
+
# verified with `transformers==4.52.1`
|
| 61 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 62 |
|
| 63 |
# Preparing model, tokenizer, and model inputs
|
|
@@ -26,13 +26,19 @@ class SinkCache(Cache):
|
|
| 26 |
|
| 27 |
Parameters:
|
| 28 |
window_length (`int`):
|
| 29 |
-
The
|
| 30 |
num_sink_tokens (`int`):
|
| 31 |
The number of sink tokens. See the original paper for more information.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
"""
|
| 33 |
|
| 34 |
-
def __init__(self, window_length: int, num_sink_tokens: int) -> None:
|
| 35 |
super().__init__()
|
|
|
|
|
|
|
| 36 |
self.key_cache: List[torch.Tensor] = []
|
| 37 |
self.value_cache: List[torch.Tensor] = []
|
| 38 |
self.window_length = window_length
|
|
@@ -40,6 +46,7 @@ class SinkCache(Cache):
|
|
| 40 |
self.cos_sin_rerotation_cache = {}
|
| 41 |
self._cos_cache = None
|
| 42 |
self._sin_cache = None
|
|
|
|
| 43 |
|
| 44 |
@staticmethod
|
| 45 |
def _rotate_half(x):
|
|
@@ -50,14 +57,18 @@ class SinkCache(Cache):
|
|
| 50 |
def _apply_key_rotary_pos_emb(
|
| 51 |
self, key_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
|
| 52 |
) -> torch.Tensor:
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
rotated_key_states = (key_states * cos) + (self._rotate_half(key_states) * sin)
|
| 54 |
-
return rotated_key_states
|
| 55 |
|
| 56 |
def _get_rerotation_cos_sin(
|
| 57 |
self, key_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
|
| 58 |
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 59 |
if key_states.shape[-2] not in self.cos_sin_rerotation_cache:
|
| 60 |
-
#
|
| 61 |
cos = cos.to(torch.float32)
|
| 62 |
sin = sin.to(torch.float32)
|
| 63 |
|
|
@@ -68,10 +79,16 @@ class SinkCache(Cache):
|
|
| 68 |
shifted_sin = sin[self.num_sink_tokens : -key_states.shape[-2]]
|
| 69 |
rerotation_cos = original_cos * shifted_cos + original_sin * shifted_sin
|
| 70 |
rerotation_sin = -original_sin * shifted_cos + original_cos * shifted_sin
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
|
| 72 |
self.cos_sin_rerotation_cache[key_states.shape[-2]] = (
|
| 73 |
-
rerotation_cos.
|
| 74 |
-
rerotation_sin.
|
| 75 |
)
|
| 76 |
return self.cos_sin_rerotation_cache[key_states.shape[-2]]
|
| 77 |
|
|
@@ -140,19 +157,29 @@ class SinkCache(Cache):
|
|
| 140 |
self.key_cache.append(key_states)
|
| 141 |
self.value_cache.append(value_states)
|
| 142 |
|
| 143 |
-
elif key_states.shape[-2] + self.get_seq_length(layer_idx) < self.window_length:
|
| 144 |
# Growing cache
|
| 145 |
self.key_cache[layer_idx] = torch.cat([self.key_cache[layer_idx], key_states], dim=-2)
|
| 146 |
self.value_cache[layer_idx] = torch.cat([self.value_cache[layer_idx], value_states], dim=-2)
|
| 147 |
|
| 148 |
else:
|
| 149 |
# Shifting cache
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 150 |
keys_to_keep = self.key_cache[layer_idx][
|
| 151 |
:, :, -self.window_length + self.num_sink_tokens + key_states.shape[-2] :
|
| 152 |
]
|
| 153 |
|
| 154 |
# On RoPE models, we need to recompute the Key rotation as the tokens are shifted
|
| 155 |
-
if using_rope:
|
| 156 |
rerotation_cos, rerotation_sin = self._get_rerotation_cos_sin(
|
| 157 |
key_states, self._cos_cache[: self.window_length], self._sin_cache[: self.window_length]
|
| 158 |
)
|
|
@@ -222,7 +249,9 @@ def generate(model, window_length=256, num_sink_tokens=4, **kwargs):
|
|
| 222 |
# 2.a. prepare the cache, if it was not passed.
|
| 223 |
past_key_values = kwargs.pop("past_key_values", None)
|
| 224 |
if past_key_values is None:
|
| 225 |
-
past_key_values = SinkCache(
|
|
|
|
|
|
|
| 226 |
elif not isinstance(past_key_values, SinkCache):
|
| 227 |
raise ValueError(f"`past_key_values` must be a `SinkCache` instance, got a {type(past_key_values)} instance")
|
| 228 |
|
|
|
|
| 26 |
|
| 27 |
Parameters:
|
| 28 |
window_length (`int`):
|
| 29 |
+
The total cache capacity, including sink tokens.
|
| 30 |
num_sink_tokens (`int`):
|
| 31 |
The number of sink tokens. See the original paper for more information.
|
| 32 |
+
position_mode (`str`, *optional*, defaults to `"auto"`):
|
| 33 |
+
`"absolute"` keeps retained RoPE keys at their original phases,
|
| 34 |
+
`"local"` rerotates them after each shift, and `"auto"` infers the
|
| 35 |
+
mode at the first eviction from `cache_position`.
|
| 36 |
"""
|
| 37 |
|
| 38 |
+
def __init__(self, window_length: int, num_sink_tokens: int, position_mode: str = "auto") -> None:
|
| 39 |
super().__init__()
|
| 40 |
+
if position_mode not in ("auto", "absolute", "local"):
|
| 41 |
+
raise ValueError("position_mode must be 'auto', 'absolute', or 'local'")
|
| 42 |
self.key_cache: List[torch.Tensor] = []
|
| 43 |
self.value_cache: List[torch.Tensor] = []
|
| 44 |
self.window_length = window_length
|
|
|
|
| 46 |
self.cos_sin_rerotation_cache = {}
|
| 47 |
self._cos_cache = None
|
| 48 |
self._sin_cache = None
|
| 49 |
+
self._global_positions = {"auto": None, "absolute": True, "local": False}[position_mode]
|
| 50 |
|
| 51 |
@staticmethod
|
| 52 |
def _rotate_half(x):
|
|
|
|
| 57 |
def _apply_key_rotary_pos_emb(
|
| 58 |
self, key_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
|
| 59 |
) -> torch.Tensor:
|
| 60 |
+
# A retained key can be rotated hundreds of times. Keep the arithmetic in fp32;
|
| 61 |
+
# rounding the one-step rotation to bf16 before multiplying makes its norm drift.
|
| 62 |
+
key_dtype = key_states.dtype
|
| 63 |
+
key_states = key_states.to(torch.float32)
|
| 64 |
rotated_key_states = (key_states * cos) + (self._rotate_half(key_states) * sin)
|
| 65 |
+
return rotated_key_states.to(key_dtype)
|
| 66 |
|
| 67 |
def _get_rerotation_cos_sin(
|
| 68 |
self, key_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
|
| 69 |
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 70 |
if key_states.shape[-2] not in self.cos_sin_rerotation_cache:
|
| 71 |
+
# Keep the rerotation coefficients in float32 through key rotation.
|
| 72 |
cos = cos.to(torch.float32)
|
| 73 |
sin = sin.to(torch.float32)
|
| 74 |
|
|
|
|
| 79 |
shifted_sin = sin[self.num_sink_tokens : -key_states.shape[-2]]
|
| 80 |
rerotation_cos = original_cos * shifted_cos + original_sin * shifted_sin
|
| 81 |
rerotation_sin = -original_sin * shifted_cos + original_cos * shifted_sin
|
| 82 |
+
# The supplied table may already be rounded to bf16 (or include a
|
| 83 |
+
# RoPE attention scale). Its derived delta is not exactly unit length.
|
| 84 |
+
# Reusing that delta on every eviction compounds the magnitude error.
|
| 85 |
+
inv_norm = torch.rsqrt(rerotation_cos.square() + rerotation_sin.square())
|
| 86 |
+
rerotation_cos = rerotation_cos * inv_norm
|
| 87 |
+
rerotation_sin = rerotation_sin * inv_norm
|
| 88 |
|
| 89 |
self.cos_sin_rerotation_cache[key_states.shape[-2]] = (
|
| 90 |
+
rerotation_cos.unsqueeze(0),
|
| 91 |
+
rerotation_sin.unsqueeze(0),
|
| 92 |
)
|
| 93 |
return self.cos_sin_rerotation_cache[key_states.shape[-2]]
|
| 94 |
|
|
|
|
| 157 |
self.key_cache.append(key_states)
|
| 158 |
self.value_cache.append(value_states)
|
| 159 |
|
| 160 |
+
elif key_states.shape[-2] + self.get_seq_length(layer_idx) <= self.window_length:
|
| 161 |
# Growing cache
|
| 162 |
self.key_cache[layer_idx] = torch.cat([self.key_cache[layer_idx], key_states], dim=-2)
|
| 163 |
self.value_cache[layer_idx] = torch.cat([self.value_cache[layer_idx], value_states], dim=-2)
|
| 164 |
|
| 165 |
else:
|
| 166 |
# Shifting cache
|
| 167 |
+
# GenerationMixin advances cache_position globally. With absolute
|
| 168 |
+
# RoPE positions, retained keys already have the correct phase;
|
| 169 |
+
# rerotating them would put them in a different coordinate system
|
| 170 |
+
# from the new query. Explicit bounded-local position streams still
|
| 171 |
+
# need rerotation after each eviction.
|
| 172 |
+
if layer_idx == 0 and self._global_positions is None:
|
| 173 |
+
cache_position = cache_kwargs.get("cache_position")
|
| 174 |
+
self._global_positions = bool(
|
| 175 |
+
cache_position is not None and cache_position[-1].item() >= self.window_length
|
| 176 |
+
)
|
| 177 |
keys_to_keep = self.key_cache[layer_idx][
|
| 178 |
:, :, -self.window_length + self.num_sink_tokens + key_states.shape[-2] :
|
| 179 |
]
|
| 180 |
|
| 181 |
# On RoPE models, we need to recompute the Key rotation as the tokens are shifted
|
| 182 |
+
if using_rope and not self._global_positions:
|
| 183 |
rerotation_cos, rerotation_sin = self._get_rerotation_cos_sin(
|
| 184 |
key_states, self._cos_cache[: self.window_length], self._sin_cache[: self.window_length]
|
| 185 |
)
|
|
|
|
| 249 |
# 2.a. prepare the cache, if it was not passed.
|
| 250 |
past_key_values = kwargs.pop("past_key_values", None)
|
| 251 |
if past_key_values is None:
|
| 252 |
+
past_key_values = SinkCache(
|
| 253 |
+
window_length=window_length, num_sink_tokens=num_sink_tokens, position_mode="absolute"
|
| 254 |
+
)
|
| 255 |
elif not isinstance(past_key_values, SinkCache):
|
| 256 |
raise ValueError(f"`past_key_values` must be a `SinkCache` instance, got a {type(past_key_values)} instance")
|
| 257 |
|
|
@@ -0,0 +1,243 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from transformers import Qwen2Config, Qwen2ForCausalLM
|
| 3 |
+
|
| 4 |
+
from custom_generate.generate import SinkCache, generate
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def test_bf16_rerotation_is_unit_length_and_keeps_cache_dtype():
|
| 8 |
+
window_length = 1024
|
| 9 |
+
head_dim = 128
|
| 10 |
+
positions = torch.arange(window_length, dtype=torch.float32)
|
| 11 |
+
inv_freq = 1_000_000.0 ** (-torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
|
| 12 |
+
angles = positions[:, None] * inv_freq[None, :]
|
| 13 |
+
cos = torch.cat((angles.cos(), angles.cos()), dim=-1).to(torch.bfloat16)
|
| 14 |
+
sin = torch.cat((angles.sin(), angles.sin()), dim=-1).to(torch.bfloat16)
|
| 15 |
+
|
| 16 |
+
cache = SinkCache(window_length=window_length, num_sink_tokens=4)
|
| 17 |
+
keys = torch.ones((1, 2, 1, head_dim), dtype=torch.bfloat16)
|
| 18 |
+
rerotation_cos, rerotation_sin = cache._get_rerotation_cos_sin(keys, cos, sin)
|
| 19 |
+
|
| 20 |
+
assert rerotation_cos.dtype == torch.float32
|
| 21 |
+
assert rerotation_sin.dtype == torch.float32
|
| 22 |
+
squared_norm = rerotation_cos.square() + rerotation_sin.square()
|
| 23 |
+
torch.testing.assert_close(squared_norm, torch.ones_like(squared_norm), atol=5e-7, rtol=0)
|
| 24 |
+
|
| 25 |
+
shifted = cache._apply_key_rotary_pos_emb(
|
| 26 |
+
keys.expand(1, 2, window_length - 5, head_dim), rerotation_cos, rerotation_sin
|
| 27 |
+
)
|
| 28 |
+
assert shifted.dtype == torch.bfloat16
|
| 29 |
+
assert torch.isfinite(shifted).all()
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def test_repeated_shift_preserves_a_key_close_to_direct_rope():
|
| 33 |
+
# A resident key walks from slot 512 to slot 4. The reference rotates its
|
| 34 |
+
# untouched pre-RoPE key directly at the final slot, without repeated writes.
|
| 35 |
+
torch.manual_seed(7)
|
| 36 |
+
capacity, sinks, head_dim = 1028, 4, 128
|
| 37 |
+
positions = torch.arange(capacity, dtype=torch.float32)
|
| 38 |
+
inv_freq = 1_000_000.0 ** (-torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
|
| 39 |
+
angles = positions[:, None] * inv_freq[None, :]
|
| 40 |
+
cos = torch.cat((angles.cos(), angles.cos()), dim=-1).to(torch.bfloat16)
|
| 41 |
+
sin = torch.cat((angles.sin(), angles.sin()), dim=-1).to(torch.bfloat16)
|
| 42 |
+
raw = torch.randn(64, head_dim).to(torch.bfloat16)
|
| 43 |
+
|
| 44 |
+
def rotate_half(value):
|
| 45 |
+
half = value.shape[-1] // 2
|
| 46 |
+
return torch.cat((-value[..., half:], value[..., :half]), dim=-1)
|
| 47 |
+
|
| 48 |
+
def rotate_once(value, position):
|
| 49 |
+
return (value * cos[position] + rotate_half(value) * sin[position]).to(torch.bfloat16)
|
| 50 |
+
|
| 51 |
+
cache = SinkCache(window_length=capacity, num_sink_tokens=sinks)
|
| 52 |
+
twiddle_cos, twiddle_sin = cache._get_rerotation_cos_sin(
|
| 53 |
+
torch.empty(1, 1, 1, head_dim, dtype=torch.bfloat16), cos, sin
|
| 54 |
+
)
|
| 55 |
+
current = rotate_once(raw, 512)
|
| 56 |
+
for destination in range(511, sinks - 1, -1):
|
| 57 |
+
offset = destination - sinks
|
| 58 |
+
current = cache._apply_key_rotary_pos_emb(
|
| 59 |
+
current, twiddle_cos[0, offset], twiddle_sin[0, offset]
|
| 60 |
+
)
|
| 61 |
+
direct = rotate_once(raw, sinks).float()
|
| 62 |
+
relative_error = ((current.float() - direct).norm(dim=-1) / direct.norm(dim=-1)).mean()
|
| 63 |
+
norm_ratio = (current.float().norm(dim=-1) / direct.norm(dim=-1)).mean()
|
| 64 |
+
assert relative_error < 0.16
|
| 65 |
+
assert 0.95 < norm_ratio < 1.05
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def test_update_applies_rope_shift_to_retained_keys():
|
| 69 |
+
cache = SinkCache(window_length=8, num_sink_tokens=1)
|
| 70 |
+
positions = torch.arange(8, dtype=torch.float32)
|
| 71 |
+
angles = positions[:, None] * torch.tensor([1.0, 0.1])
|
| 72 |
+
cos = torch.cat((angles.cos(), angles.cos()), dim=-1).to(torch.bfloat16)
|
| 73 |
+
sin = torch.cat((angles.sin(), angles.sin()), dim=-1).to(torch.bfloat16)
|
| 74 |
+
old_keys = torch.arange(32, dtype=torch.float32).reshape(1, 1, 8, 4).to(torch.bfloat16)
|
| 75 |
+
values = torch.zeros_like(old_keys)
|
| 76 |
+
cache.update(old_keys, values, 0, {"cos": cos, "sin": sin})
|
| 77 |
+
old_window = cache.key_cache[0][..., 2:, :].clone()
|
| 78 |
+
|
| 79 |
+
new_key = torch.ones((1, 1, 1, 4), dtype=torch.bfloat16)
|
| 80 |
+
cache.update(new_key, torch.zeros_like(new_key), 0, {"cos": cos, "sin": sin})
|
| 81 |
+
|
| 82 |
+
assert cache.get_seq_length() == 8
|
| 83 |
+
assert cache.key_cache[0].dtype == torch.bfloat16
|
| 84 |
+
assert torch.equal(cache.key_cache[0][..., :1, :], old_keys[..., :1, :])
|
| 85 |
+
assert torch.equal(cache.key_cache[0][..., -1:, :], new_key)
|
| 86 |
+
assert not torch.equal(cache.key_cache[0][..., 1:-1, :], old_window)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def test_absolute_positions_keep_retained_keys_in_original_rope_frame():
|
| 90 |
+
cache = SinkCache(window_length=8, num_sink_tokens=1)
|
| 91 |
+
old_keys = torch.arange(32, dtype=torch.float32).reshape(1, 1, 8, 4).to(torch.bfloat16)
|
| 92 |
+
values = torch.zeros_like(old_keys)
|
| 93 |
+
cos = torch.arange(12, dtype=torch.float32)[:, None].expand(12, 4).cos().to(torch.bfloat16)
|
| 94 |
+
sin = torch.arange(12, dtype=torch.float32)[:, None].expand(12, 4).sin().to(torch.bfloat16)
|
| 95 |
+
cache.update(old_keys, values, 0, {"cos": cos, "sin": sin, "cache_position": torch.arange(8)})
|
| 96 |
+
|
| 97 |
+
new_key = torch.ones((1, 1, 1, 4), dtype=torch.bfloat16)
|
| 98 |
+
cache.update(new_key, torch.zeros_like(new_key), 0,
|
| 99 |
+
{"cos": cos, "sin": sin, "cache_position": torch.tensor([8])})
|
| 100 |
+
|
| 101 |
+
assert cache._global_positions is True
|
| 102 |
+
assert cache.get_seq_length() == 8
|
| 103 |
+
assert torch.equal(cache.key_cache[0][..., :1, :], old_keys[..., :1, :])
|
| 104 |
+
assert torch.equal(cache.key_cache[0][..., 1:-1, :], old_keys[..., 2:, :])
|
| 105 |
+
assert torch.equal(cache.key_cache[0][..., -1:, :], new_key)
|
| 106 |
+
assert not cache.cos_sin_rerotation_cache
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def test_explicit_absolute_mode_does_not_guess_from_position_index():
|
| 110 |
+
cache = SinkCache(window_length=8, num_sink_tokens=1, position_mode="absolute")
|
| 111 |
+
old_keys = torch.arange(32, dtype=torch.float32).reshape(1, 1, 8, 4)
|
| 112 |
+
cache.update(old_keys, old_keys, 0)
|
| 113 |
+
new_key = torch.ones((1, 1, 1, 4))
|
| 114 |
+
cos = torch.ones((8, 4))
|
| 115 |
+
sin = torch.zeros((8, 4))
|
| 116 |
+
cache.update(new_key, new_key, 0,
|
| 117 |
+
{"cos": cos, "sin": sin, "cache_position": torch.tensor([7])})
|
| 118 |
+
assert cache._global_positions is True
|
| 119 |
+
assert torch.equal(cache.key_cache[0][..., 1:-1, :], old_keys[..., 2:, :])
|
| 120 |
+
assert not cache.cos_sin_rerotation_cache
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
def test_exactly_filling_capacity_does_not_evict():
|
| 124 |
+
cache = SinkCache(window_length=8, num_sink_tokens=1)
|
| 125 |
+
old_keys = torch.arange(28, dtype=torch.float32).reshape(1, 1, 7, 4)
|
| 126 |
+
cache.update(old_keys, old_keys, 0)
|
| 127 |
+
new_key = torch.ones((1, 1, 1, 4))
|
| 128 |
+
cache.update(new_key, new_key, 0)
|
| 129 |
+
assert cache.get_seq_length() == 8
|
| 130 |
+
assert torch.equal(cache.key_cache[0][..., :7, :], old_keys)
|
| 131 |
+
assert torch.equal(cache.key_cache[0][..., 7:, :], new_key)
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def test_qwen2_native_forward_supplies_rope_and_runs_shifts():
|
| 135 |
+
config = Qwen2Config(
|
| 136 |
+
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
|
| 137 |
+
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
|
| 138 |
+
rope_theta=1_000_000.0,
|
| 139 |
+
)
|
| 140 |
+
config._attn_implementation = "eager"
|
| 141 |
+
model = Qwen2ForCausalLM(config).eval()
|
| 142 |
+
cache = SinkCache(window_length=8, num_sink_tokens=1)
|
| 143 |
+
with torch.inference_mode():
|
| 144 |
+
model(
|
| 145 |
+
input_ids=torch.arange(8).unsqueeze(0), past_key_values=cache,
|
| 146 |
+
cache_position=torch.arange(8), use_cache=True,
|
| 147 |
+
)
|
| 148 |
+
assert not cache.cos_sin_rerotation_cache
|
| 149 |
+
for token in range(8, 24):
|
| 150 |
+
model(
|
| 151 |
+
input_ids=torch.tensor([[token % 64]]), past_key_values=cache,
|
| 152 |
+
cache_position=torch.tensor([7]), use_cache=True,
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
assert cache.get_seq_length() == 8
|
| 156 |
+
assert 1 in cache.cos_sin_rerotation_cache
|
| 157 |
+
assert len(cache.key_cache) == config.num_hidden_layers
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
def test_qwen2_generate_uses_absolute_positions_without_rerotation():
|
| 161 |
+
config = Qwen2Config(
|
| 162 |
+
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
|
| 163 |
+
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
|
| 164 |
+
rope_theta=1_000_000.0,
|
| 165 |
+
)
|
| 166 |
+
config._attn_implementation = "eager"
|
| 167 |
+
model = Qwen2ForCausalLM(config).eval()
|
| 168 |
+
cache = SinkCache(window_length=8, num_sink_tokens=1)
|
| 169 |
+
with torch.inference_mode():
|
| 170 |
+
outputs = model.generate(
|
| 171 |
+
input_ids=torch.arange(8).unsqueeze(0), max_new_tokens=5,
|
| 172 |
+
do_sample=False, past_key_values=cache, use_cache=True,
|
| 173 |
+
)
|
| 174 |
+
assert outputs.shape == (1, 13)
|
| 175 |
+
assert cache.get_seq_length() == 8
|
| 176 |
+
assert cache._global_positions is True
|
| 177 |
+
assert not cache.cos_sin_rerotation_cache
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
def test_custom_generate_wrapper_selects_absolute_mode():
|
| 181 |
+
config = Qwen2Config(
|
| 182 |
+
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
|
| 183 |
+
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
|
| 184 |
+
rope_theta=1_000_000.0,
|
| 185 |
+
)
|
| 186 |
+
config._attn_implementation = "eager"
|
| 187 |
+
model = Qwen2ForCausalLM(config).eval()
|
| 188 |
+
with torch.inference_mode():
|
| 189 |
+
outputs = generate(
|
| 190 |
+
model, window_length=8, num_sink_tokens=1,
|
| 191 |
+
input_ids=torch.arange(8).unsqueeze(0), max_new_tokens=5,
|
| 192 |
+
do_sample=False, return_dict_in_generate=True,
|
| 193 |
+
)
|
| 194 |
+
cache = outputs.past_key_values
|
| 195 |
+
assert isinstance(cache, SinkCache)
|
| 196 |
+
assert cache._global_positions is True
|
| 197 |
+
assert not cache.cos_sin_rerotation_cache
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
def test_explicit_absolute_positions_match_qwen2_default_forward():
|
| 201 |
+
config = Qwen2Config(
|
| 202 |
+
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
|
| 203 |
+
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
|
| 204 |
+
rope_theta=1_000_000.0,
|
| 205 |
+
)
|
| 206 |
+
config._attn_implementation = "eager"
|
| 207 |
+
model = Qwen2ForCausalLM(config).eval()
|
| 208 |
+
caches = [SinkCache(window_length=8, num_sink_tokens=1) for _ in range(2)]
|
| 209 |
+
with torch.inference_mode():
|
| 210 |
+
for cache in caches:
|
| 211 |
+
model(input_ids=torch.arange(8).unsqueeze(0), past_key_values=cache,
|
| 212 |
+
cache_position=torch.arange(8), use_cache=True)
|
| 213 |
+
explicit = model(input_ids=torch.tensor([[8]]), past_key_values=caches[0],
|
| 214 |
+
cache_position=torch.tensor([8]), position_ids=torch.tensor([[8]]),
|
| 215 |
+
use_cache=True).logits
|
| 216 |
+
default = model(input_ids=torch.tensor([[8]]), past_key_values=caches[1],
|
| 217 |
+
cache_position=torch.tensor([8]), use_cache=True).logits
|
| 218 |
+
torch.testing.assert_close(explicit, default, atol=0, rtol=0)
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def test_explicit_absolute_and_auto_modes_match_on_standard_stream():
|
| 222 |
+
config = Qwen2Config(
|
| 223 |
+
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
|
| 224 |
+
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
|
| 225 |
+
rope_theta=1_000_000.0,
|
| 226 |
+
)
|
| 227 |
+
config._attn_implementation = "eager"
|
| 228 |
+
model = Qwen2ForCausalLM(config).eval()
|
| 229 |
+
caches = [SinkCache(8, 1, "auto"), SinkCache(8, 1, "absolute")]
|
| 230 |
+
with torch.inference_mode():
|
| 231 |
+
for cache in caches:
|
| 232 |
+
model(input_ids=torch.arange(8).unsqueeze(0), past_key_values=cache,
|
| 233 |
+
cache_position=torch.arange(8), use_cache=True)
|
| 234 |
+
for token in range(8, 20):
|
| 235 |
+
logits = []
|
| 236 |
+
for cache in caches:
|
| 237 |
+
logits.append(model(input_ids=torch.tensor([[token % 64]]),
|
| 238 |
+
past_key_values=cache, cache_position=torch.tensor([token]),
|
| 239 |
+
position_ids=torch.tensor([[token]]), use_cache=True).logits)
|
| 240 |
+
torch.testing.assert_close(logits[0], logits[1], atol=0, rtol=0)
|
| 241 |
+
assert caches[0]._global_positions is True
|
| 242 |
+
assert not caches[0].cos_sin_rerotation_cache
|
| 243 |
+
assert not caches[1].cos_sin_rerotation_cache
|