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