jholl-hugging commited on
Commit
5382829
·
verified ·
1 Parent(s): ea681c7

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 CHANGED
@@ -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 matches the `SinkCache` class present in `transformers<4.53.0`.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15
 
16
  ![Sink Cache diagram from the original paper](https://arxiv.org/html/2309.17453v4/x1.png)
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 length of the context window.
 
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
- # requires `transformers>=4.52.0`
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
  ![Sink Cache diagram from the original paper](https://arxiv.org/html/2309.17453v4/x1.png)
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
custom_generate/generate.py CHANGED
@@ -26,13 +26,19 @@ class SinkCache(Cache):
26
 
27
  Parameters:
28
  window_length (`int`):
29
- The length of the context window.
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
- # Upcast to float32 temporarily for better accuracy
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.to(key_states.dtype).unsqueeze(0),
74
- rerotation_sin.to(key_states.dtype).unsqueeze(0),
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(window_length=window_length, num_sink_tokens=num_sink_tokens)
 
 
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
 
tests/test_sink_cache_precision.py ADDED
@@ -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