Q-TensorFormer / tests /test_kv_cache.py
Premchandyadav369
Transform Q-TensorFormer into an Information-Value Adaptive Resource Allocation Architecture
eaeea8f
Raw History Blame Contribute Delete
2.78 kB
"""
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!")