# vLLM中的KV Cache实现原理
from vllm.attention import Attention
from vllm.model_executor.layers.attention import PagedAttention

class KVCache:
    def __init__(self, num_layers, num_heads, head_dim, max_seq_len):
        # vLLM使用分页内存管理KV Cache
        self.k_cache = torch.zeros(
            num_layers, max_seq_len, num_heads, head_dim,
            dtype=torch.float16, device='cuda'
        )
        self.v_cache = torch.zeros(
            num_layers, max_seq_len, num_heads, head_dim,
            dtype=torch.float16, device='cuda'
        )
    
    def update_cache(self, layer_id, new_k, new_v, seq_pos):
        """更新指定位置的KV缓存"""
        self.k_cache[layer_id, seq_pos] = new_k
        self.v_cache[layer_id, seq_pos] = new_v
    
    def get_cached_kv(self, layer_id, seq_len):
        """获取缓存的KV值"""
        return (
            self.k_cache[layer_id, :seq_len],
            self.v_cache[layer_id, :seq_len]
        )
