File size: 2,776 Bytes
eaeea8f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
"""
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!")