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
Download tests/test_sink_cache_precision.py from transformers-community/sink_cache: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/transformers-community/sink_cache/resolve/refs%2Fpr%2F2/tests/test_sink_cache_precision.py
- Command line
-
hf download hf://transformers-community/sink_cache@refs/pr/2/tests/test_sink_cache_precision.py
-
curl -L -o test_sink_cache_precision.py https://huggingface.co/transformers-community/sink_cache/resolve/refs%2Fpr%2F2/tests/test_sink_cache_precision.py
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 | |