""" Tests for Adaptive KV Cache Module. """ import sys import os sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import torch import pytest from src.kv_cache import AdaptiveKVCache, KVPrecision, QuantizedKVTensor def test_quantized_kv_tensor(): tensor = torch.randn(2, 4, 8, 32) # FP16 q_fp16 = QuantizedKVTensor(tensor, KVPrecision.FP16) rec_fp16 = q_fp16.dequantize() assert torch.allclose(tensor, rec_fp16, atol=1e-2) # INT8 q_int8 = QuantizedKVTensor(tensor, KVPrecision.INT8) rec_int8 = q_int8.dequantize() cos_sim_int8 = torch.cosine_similarity(tensor.flatten(), rec_int8.flatten(), dim=0) assert cos_sim_int8 > 0.99, f"INT8 cosine similarity {cos_sim_int8} too low" # INT4 q_int4 = QuantizedKVTensor(tensor, KVPrecision.INT4) rec_int4 = q_int4.dequantize() cos_sim_int4 = torch.cosine_similarity(tensor.flatten(), rec_int4.flatten(), dim=0) assert cos_sim_int4 > 0.90, f"INT4 cosine similarity {cos_sim_int4} too low" print("✓ test_quantized_kv_tensor passed") def test_adaptive_kv_cache_append_and_evict(): cache = AdaptiveKVCache(max_capacity=16, default_precision=KVPrecision.FP16, window_size=4) B, H, D = 1, 2, 16 # Append 10 tokens k1 = torch.randn(B, H, 10, D) v1 = torch.randn(B, H, 10, D) out_k1, out_v1 = cache.update(k1, v1) assert cache.seq_len == 10 assert out_k1.shape[-2] == 10 # Append 10 more tokens (exceeds max_capacity 16 -> should trigger eviction) k2 = torch.randn(B, H, 10, D) v2 = torch.randn(B, H, 10, D) out_k2, out_v2 = cache.update(k2, v2) assert cache.seq_len == 16, f"Expected cache seq_len 16, got {cache.seq_len}" assert cache.evicted_tokens_count == 4 assert cache.current_mb > 0.0 print("✓ test_adaptive_kv_cache_append_and_evict passed") def test_adaptive_kv_cache_precision_switch(): cache = AdaptiveKVCache(max_capacity=32, default_precision=KVPrecision.FP16) k = torch.randn(1, 2, 8, 16) v = torch.randn(1, 2, 8, 16) cache.update(k, v) bytes_fp16 = cache.current_bytes # Switch to INT8 cache.set_precision(KVPrecision.INT8) bytes_int8 = cache.current_bytes assert bytes_int8 < bytes_fp16, f"INT8 ({bytes_int8}) should be smaller than FP16 ({bytes_fp16})" # Switch to INT4 cache.set_precision(KVPrecision.INT4) bytes_int4 = cache.current_bytes assert bytes_int4 < bytes_int8, f"INT4 ({bytes_int4}) should be smaller than INT8 ({bytes_int8})" print("✓ test_adaptive_kv_cache_precision_switch passed") if __name__ == "__main__": test_quantized_kv_tensor() test_adaptive_kv_cache_append_and_evict() test_adaptive_kv_cache_precision_switch() print("All KV Cache tests passed!")