sink_cache / tests /test_sink_cache_precision.py
jholl-hugging's picture
Fix SinkCache RoPE coordinates for absolute generation
5382829 verified
Raw History Blame
11.4 kB
import torch
from transformers import Qwen2Config, Qwen2ForCausalLM
from custom_generate.generate import SinkCache, generate
def test_bf16_rerotation_is_unit_length_and_keeps_cache_dtype():
window_length = 1024
head_dim = 128
positions = torch.arange(window_length, dtype=torch.float32)
inv_freq = 1_000_000.0 ** (-torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
angles = positions[:, None] * inv_freq[None, :]
cos = torch.cat((angles.cos(), angles.cos()), dim=-1).to(torch.bfloat16)
sin = torch.cat((angles.sin(), angles.sin()), dim=-1).to(torch.bfloat16)
cache = SinkCache(window_length=window_length, num_sink_tokens=4)
keys = torch.ones((1, 2, 1, head_dim), dtype=torch.bfloat16)
rerotation_cos, rerotation_sin = cache._get_rerotation_cos_sin(keys, cos, sin)
assert rerotation_cos.dtype == torch.float32
assert rerotation_sin.dtype == torch.float32
squared_norm = rerotation_cos.square() + rerotation_sin.square()
torch.testing.assert_close(squared_norm, torch.ones_like(squared_norm), atol=5e-7, rtol=0)
shifted = cache._apply_key_rotary_pos_emb(
keys.expand(1, 2, window_length - 5, head_dim), rerotation_cos, rerotation_sin
)
assert shifted.dtype == torch.bfloat16
assert torch.isfinite(shifted).all()
def test_repeated_shift_preserves_a_key_close_to_direct_rope():
# A resident key walks from slot 512 to slot 4. The reference rotates its
# untouched pre-RoPE key directly at the final slot, without repeated writes.
torch.manual_seed(7)
capacity, sinks, head_dim = 1028, 4, 128
positions = torch.arange(capacity, dtype=torch.float32)
inv_freq = 1_000_000.0 ** (-torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
angles = positions[:, None] * inv_freq[None, :]
cos = torch.cat((angles.cos(), angles.cos()), dim=-1).to(torch.bfloat16)
sin = torch.cat((angles.sin(), angles.sin()), dim=-1).to(torch.bfloat16)
raw = torch.randn(64, head_dim).to(torch.bfloat16)
def rotate_half(value):
half = value.shape[-1] // 2
return torch.cat((-value[..., half:], value[..., :half]), dim=-1)
def rotate_once(value, position):
return (value * cos[position] + rotate_half(value) * sin[position]).to(torch.bfloat16)
cache = SinkCache(window_length=capacity, num_sink_tokens=sinks)
twiddle_cos, twiddle_sin = cache._get_rerotation_cos_sin(
torch.empty(1, 1, 1, head_dim, dtype=torch.bfloat16), cos, sin
)
current = rotate_once(raw, 512)
for destination in range(511, sinks - 1, -1):
offset = destination - sinks
current = cache._apply_key_rotary_pos_emb(
current, twiddle_cos[0, offset], twiddle_sin[0, offset]
)
direct = rotate_once(raw, sinks).float()
relative_error = ((current.float() - direct).norm(dim=-1) / direct.norm(dim=-1)).mean()
norm_ratio = (current.float().norm(dim=-1) / direct.norm(dim=-1)).mean()
assert relative_error < 0.16
assert 0.95 < norm_ratio < 1.05
def test_update_applies_rope_shift_to_retained_keys():
cache = SinkCache(window_length=8, num_sink_tokens=1)
positions = torch.arange(8, dtype=torch.float32)
angles = positions[:, None] * torch.tensor([1.0, 0.1])
cos = torch.cat((angles.cos(), angles.cos()), dim=-1).to(torch.bfloat16)
sin = torch.cat((angles.sin(), angles.sin()), dim=-1).to(torch.bfloat16)
old_keys = torch.arange(32, dtype=torch.float32).reshape(1, 1, 8, 4).to(torch.bfloat16)
values = torch.zeros_like(old_keys)
cache.update(old_keys, values, 0, {"cos": cos, "sin": sin})
old_window = cache.key_cache[0][..., 2:, :].clone()
new_key = torch.ones((1, 1, 1, 4), dtype=torch.bfloat16)
cache.update(new_key, torch.zeros_like(new_key), 0, {"cos": cos, "sin": sin})
assert cache.get_seq_length() == 8
assert cache.key_cache[0].dtype == torch.bfloat16
assert torch.equal(cache.key_cache[0][..., :1, :], old_keys[..., :1, :])
assert torch.equal(cache.key_cache[0][..., -1:, :], new_key)
assert not torch.equal(cache.key_cache[0][..., 1:-1, :], old_window)
def test_absolute_positions_keep_retained_keys_in_original_rope_frame():
cache = SinkCache(window_length=8, num_sink_tokens=1)
old_keys = torch.arange(32, dtype=torch.float32).reshape(1, 1, 8, 4).to(torch.bfloat16)
values = torch.zeros_like(old_keys)
cos = torch.arange(12, dtype=torch.float32)[:, None].expand(12, 4).cos().to(torch.bfloat16)
sin = torch.arange(12, dtype=torch.float32)[:, None].expand(12, 4).sin().to(torch.bfloat16)
cache.update(old_keys, values, 0, {"cos": cos, "sin": sin, "cache_position": torch.arange(8)})
new_key = torch.ones((1, 1, 1, 4), dtype=torch.bfloat16)
cache.update(new_key, torch.zeros_like(new_key), 0,
{"cos": cos, "sin": sin, "cache_position": torch.tensor([8])})
assert cache._global_positions is True
assert cache.get_seq_length() == 8
assert torch.equal(cache.key_cache[0][..., :1, :], old_keys[..., :1, :])
assert torch.equal(cache.key_cache[0][..., 1:-1, :], old_keys[..., 2:, :])
assert torch.equal(cache.key_cache[0][..., -1:, :], new_key)
assert not cache.cos_sin_rerotation_cache
def test_explicit_absolute_mode_does_not_guess_from_position_index():
cache = SinkCache(window_length=8, num_sink_tokens=1, position_mode="absolute")
old_keys = torch.arange(32, dtype=torch.float32).reshape(1, 1, 8, 4)
cache.update(old_keys, old_keys, 0)
new_key = torch.ones((1, 1, 1, 4))
cos = torch.ones((8, 4))
sin = torch.zeros((8, 4))
cache.update(new_key, new_key, 0,
{"cos": cos, "sin": sin, "cache_position": torch.tensor([7])})
assert cache._global_positions is True
assert torch.equal(cache.key_cache[0][..., 1:-1, :], old_keys[..., 2:, :])
assert not cache.cos_sin_rerotation_cache
def test_exactly_filling_capacity_does_not_evict():
cache = SinkCache(window_length=8, num_sink_tokens=1)
old_keys = torch.arange(28, dtype=torch.float32).reshape(1, 1, 7, 4)
cache.update(old_keys, old_keys, 0)
new_key = torch.ones((1, 1, 1, 4))
cache.update(new_key, new_key, 0)
assert cache.get_seq_length() == 8
assert torch.equal(cache.key_cache[0][..., :7, :], old_keys)
assert torch.equal(cache.key_cache[0][..., 7:, :], new_key)
def test_qwen2_native_forward_supplies_rope_and_runs_shifts():
config = Qwen2Config(
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
rope_theta=1_000_000.0,
)
config._attn_implementation = "eager"
model = Qwen2ForCausalLM(config).eval()
cache = SinkCache(window_length=8, num_sink_tokens=1)
with torch.inference_mode():
model(
input_ids=torch.arange(8).unsqueeze(0), past_key_values=cache,
cache_position=torch.arange(8), use_cache=True,
)
assert not cache.cos_sin_rerotation_cache
for token in range(8, 24):
model(
input_ids=torch.tensor([[token % 64]]), past_key_values=cache,
cache_position=torch.tensor([7]), use_cache=True,
)
assert cache.get_seq_length() == 8
assert 1 in cache.cos_sin_rerotation_cache
assert len(cache.key_cache) == config.num_hidden_layers
def test_qwen2_generate_uses_absolute_positions_without_rerotation():
config = Qwen2Config(
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
rope_theta=1_000_000.0,
)
config._attn_implementation = "eager"
model = Qwen2ForCausalLM(config).eval()
cache = SinkCache(window_length=8, num_sink_tokens=1)
with torch.inference_mode():
outputs = model.generate(
input_ids=torch.arange(8).unsqueeze(0), max_new_tokens=5,
do_sample=False, past_key_values=cache, use_cache=True,
)
assert outputs.shape == (1, 13)
assert cache.get_seq_length() == 8
assert cache._global_positions is True
assert not cache.cos_sin_rerotation_cache
def test_custom_generate_wrapper_selects_absolute_mode():
config = Qwen2Config(
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
rope_theta=1_000_000.0,
)
config._attn_implementation = "eager"
model = Qwen2ForCausalLM(config).eval()
with torch.inference_mode():
outputs = generate(
model, window_length=8, num_sink_tokens=1,
input_ids=torch.arange(8).unsqueeze(0), max_new_tokens=5,
do_sample=False, return_dict_in_generate=True,
)
cache = outputs.past_key_values
assert isinstance(cache, SinkCache)
assert cache._global_positions is True
assert not cache.cos_sin_rerotation_cache
def test_explicit_absolute_positions_match_qwen2_default_forward():
config = Qwen2Config(
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
rope_theta=1_000_000.0,
)
config._attn_implementation = "eager"
model = Qwen2ForCausalLM(config).eval()
caches = [SinkCache(window_length=8, num_sink_tokens=1) for _ in range(2)]
with torch.inference_mode():
for cache in caches:
model(input_ids=torch.arange(8).unsqueeze(0), past_key_values=cache,
cache_position=torch.arange(8), use_cache=True)
explicit = model(input_ids=torch.tensor([[8]]), past_key_values=caches[0],
cache_position=torch.tensor([8]), position_ids=torch.tensor([[8]]),
use_cache=True).logits
default = model(input_ids=torch.tensor([[8]]), past_key_values=caches[1],
cache_position=torch.tensor([8]), use_cache=True).logits
torch.testing.assert_close(explicit, default, atol=0, rtol=0)
def test_explicit_absolute_and_auto_modes_match_on_standard_stream():
config = Qwen2Config(
vocab_size=64, hidden_size=64, intermediate_size=128, num_hidden_layers=2,
num_attention_heads=4, num_key_value_heads=2, max_position_embeddings=64,
rope_theta=1_000_000.0,
)
config._attn_implementation = "eager"
model = Qwen2ForCausalLM(config).eval()
caches = [SinkCache(8, 1, "auto"), SinkCache(8, 1, "absolute")]
with torch.inference_mode():
for cache in caches:
model(input_ids=torch.arange(8).unsqueeze(0), past_key_values=cache,
cache_position=torch.arange(8), use_cache=True)
for token in range(8, 20):
logits = []
for cache in caches:
logits.append(model(input_ids=torch.tensor([[token % 64]]),
past_key_values=cache, cache_position=torch.tensor([token]),
position_ids=torch.tensor([[token]]), use_cache=True).logits)
torch.testing.assert_close(logits[0], logits[1], atol=0, rtol=0)
assert caches[0]._global_positions is True
assert not caches[0].cos_sin_rerotation_cache
assert not caches[1].cos_sin_rerotation_cache