5.KV Cache 管理 & PagedAttention
实现一个函数,计算给定模型配置(层数、头数、头维度、最大序列长度、batch size)下 KV Cache 的总显存占用量(以 FP16 存储)。¶
1.1 核心公式推导¶
在标准的多头自注意力(MHA)中,每一层都需要存储 Key 和 Value 两个张量。对于单个 token,其在某一层的 K 和 V 都是形状为 (num_heads, head_dim) 的张量。因此,每 token 每层的 KV 数据量为:

在 FP16 下,bytes_per_elem = 2。扩展到整个序列长度 max_seq_len 和 batch size batch_size,以及所有 num_layers 层,总显存占用量为:

如果使用 GQA(分组查询注意力)或 MQA(多查询注意力),只需将 num_heads 替换为实际参与 KV 存储的头数 num_kv_heads。例如,LLaMA-2 70B 中 num_kv_heads=8,num_heads=64,此时显存会大幅降低。
1.2 代码实现¶
import math
def compute_kv_cache_memory(
num_layers: int,
num_heads: int, # 或者 num_kv_heads
head_dim: int,
max_seq_len: int,
batch_size: int,
dtype_bytes: int = 2, # FP16
is_gqa: bool = False,
num_kv_heads: int = None
) -> dict:
"""
计算 KV Cache 总显存(字节),同时返回可读的格式化字符串。
支持 GQA/MQA:传入 is_gqa=True 和 num_kv_heads。
"""
if is_gqa and num_kv_heads is not None:
heads = num_kv_heads
else:
heads = num_heads
# 每 token 每层的 K 或 V 大小
per_token_kv_bytes = heads * head_dim * dtype_bytes
# 每 token 每层的 K+V 总大小
per_token_layer_bytes = 2 * per_token_kv_bytes
# 总显存 = 层数 * batch * 最大长度 * 每token每层大小
total_bytes = num_layers * batch_size * max_seq_len * per_token_layer_bytes
# 转换为可读单位
units = ["B", "KiB", "MiB", "GiB", "TiB"]
value = total_bytes
unit_idx = 0
while value >= 1024 and unit_idx < len(units) - 1:
value /= 1024.0
unit_idx += 1
return {
"total_bytes": total_bytes,
"readable": f"{value:.2f} {units[unit_idx]}",
"per_token_layer_bytes": per_token_layer_bytes,
"heads_used": heads
}
# 示例:LLaMA-2 7B (num_layers=32, num_heads=32, head_dim=128)
mem_info = compute_kv_cache_memory(32, 32, 128, 2048, 4)
print(mem_info["readable"]) # 约 4 GiB
1.3 扩展到量化缓存¶
如果 KV Cache 本身也被量化为 INT8 或 FP8,显存占用会进一步减半或更多。例如,使用 FP8 时 dtype_bytes=1。上述函数可通过改变 dtype_bytes 参数覆盖此场景。
1.4 注意事项¶
-
上述计算是理论下界,实际引擎(如 vLLM)中由于分页管理、块内对齐等因素,实际占用可能略高(通常有 <5% 的碎片开销)。
-
当使用 PagedAttention 时,显存占用受物理块大小的限制,内部碎片(最后一个块未满)会增加少量开销,但整体仍然接近理论值。
-
在分布式推理(张量并行)中,KV Cache 被切分到多个 GPU 上,每张 GPU 的显存占用为上述计算值除以并行路数(假设均匀切分)。
用 Python 实现一个简单的 KV Cache 管理器,支持初始化、按序列存储 key/value,以及按序列 ID 和位置区间读取缓存。¶
2.1 设计目标¶
我们需要一个能够管理多个序列、支持多层注意力的 KV Cache 管理器。它应具备:
-
按层、按序列存储 K 和 V。
-
支持追加写入(decode 阶段)和区间读取(prefill 后获取全部、decode 时读取历史)。
-
高效的内存管理,避免频繁分配。
2.2 基于字典的简单实现¶
import torch
from typing import Dict, Tuple, Optional
class SimpleKVCache:
def __init__(self):
# 存储结构: self.cache[layer_idx][seq_id] = (K_tensor, V_tensor)
self.cache: Dict[int, Dict[int, Tuple[torch.Tensor, torch.Tensor]]] = {}
def init_sequence(self, layer_idx: int, seq_id: int, k: torch.Tensor, v: torch.Tensor):
"""初始化或覆盖某序列在某层的 KV。k/v 形状 [num_heads, seq_len, head_dim]"""
if layer_idx not in self.cache:
self.cache[layer_idx] = {}
self.cache[layer_idx][seq_id] = (k, v)
def append(self, layer_idx: int, seq_id: int, new_k: torch.Tensor, new_v: torch.Tensor):
"""追加单个 token 的 KV。new_k 形状 [num_heads, 1, head_dim]"""
if layer_idx not in self.cache or seq_id not in self.cache[layer_idx]:
# 如果不存在则初始化
self.init_sequence(layer_idx, seq_id, new_k, new_v)
return
k, v = self.cache[layer_idx][seq_id]
self.cache[layer_idx][seq_id] = (
torch.cat([k, new_k], dim=1),
torch.cat([v, new_v], dim=1)
)
def get(self, layer_idx: int, seq_id: int, start: int = 0, end: Optional[int] = None) -> Tuple[torch.Tensor, torch.Tensor]:
"""读取区间 [start, end) 的 KV。若 end 为 None,则读到末尾。"""
k, v = self.cache[layer_idx][seq_id]
if end is None:
end = k.size(1)
return k[:, start:end, :], v[:, start:end, :]
def get_sequence_length(self, layer_idx: int, seq_id: int) -> int:
return self.cache[layer_idx][seq_id][0].size(1)
def remove_sequence(self, layer_idx: int, seq_id: int):
if layer_idx in self.cache and seq_id in self.cache[layer_idx]:
del self.cache[layer_idx][seq_id]
2.3 使用 PyTorch 张量池的优化实现¶
频繁的 torch.cat 会导致内存碎片和性能下降。实际系统中常预分配一块大的张量作为缓存池,然后通过索引写入。例如,为每层分配一个 [max_batch_size, max_seq_len, num_heads, head_dim] 的 K 和 V 张量,然后按序列 ID 写入。
class PreallocKVCache:
def __init__(self, num_layers, max_batch, max_seq_len, num_heads, head_dim):
self.num_layers = num_layers
# 为每层预分配 K 和 V 张量
self.k = torch.zeros(num_layers, max_batch, num_heads, max_seq_len, head_dim, dtype=torch.float16)
self.v = torch.zeros_like(self.k)
self.seq_len = torch.zeros(max_batch, dtype=torch.long) # 每个batch槽的当前长度
self.seq_to_slot = {} # seq_id -> batch_slot
def register_seq(self, seq_id: int) -> int:
"""为新序列分配一个batch槽位"""
slot = len(self.seq_to_slot)
self.seq_to_slot[seq_id] = slot
return slot
def store_prefill(self, layer_idx: int, seq_id: int, k: torch.Tensor, v: torch.Tensor):
"""prefill 时写入整个 prompt 的 KV"""
slot = self.seq_to_slot[seq_id]
L = k.size(2) # k shape: [1, heads, L, D]
self.k[layer_idx, slot, :, :L] = k[0]
self.v[layer_idx, slot, :, :L] = v[0]
self.seq_len[slot] = L
def append_decode(self, layer_idx: int, seq_id: int, new_k: torch.Tensor, new_v: torch.Tensor):
"""decode 时追加一个 token"""
slot = self.seq_to_slot[seq_id]
cur_len = self.seq_len[slot]
self.k[layer_idx, slot, :, cur_len:cur_len+1] = new_k[0]
self.v[layer_idx, slot, :, cur_len:cur_len+1] = new_v[0]
self.seq_len[slot] += 1
这种预分配方案避免了动态分配,且在 GPU 上直接操作,效率极高,但需要预先设定最大 batch 和最大长度。
2.4 对比与选择¶
-
简单字典:灵活性高,适合原型验证。
-
预分配张量:性能高,适合生产环境,类似 ONNX Runtime 的 IOBinding。
-
PagedAttention:终极方案,通过分页管理消除预分配浪费,vLLM 的核心。
写出 prefill 阶段为 batch 中多条序列分配并填充 KV Cache 的代码框架,假设输入 shape 为 [B, S, H],输出对应 key/value 并存入缓存。¶
3.1 Prefill 与 Decode 的区别¶
-
Prefill:一次性处理整个 prompt 序列,产生完整的 K 和 V。此时 self-attention 是双向的(或因果),需要一次性计算所有 token 的注意力。
-
Decode:每次只处理一个新 token,利用已缓存的 K、V 进行增量计算。
3.2 通用代码框架(基于 PreallocKVCache)¶
def prefill_with_cache(model, input_ids, seq_ids, cache: PreallocKVCache):
"""
input_ids: [B, S] token IDs
seq_ids: list of seq_id for each batch element
假设 model 提供 get_kv_after_forward 返回每层的 K/V
"""
B, S = input_ids.shape
# 为每个序列注册 batch slot
for sid in seq_ids:
if sid not in cache.seq_to_slot:
cache.register_seq(sid)
# 模型前向传播(这里用伪代码表示)
# hidden_states, kv_per_layer = model.forward(input_ids, return_kv=True)
# kv_per_layer: list of (k, v),每个形状 [B, num_heads, S, head_dim]
for layer_idx, (k, v) in enumerate(kv_per_layer):
for i, sid in enumerate(seq_ids):
# 取出该序列在该层的 K、V
k_seq = k[i:i+1] # [1, num_heads, S, head_dim]
v_seq = v[i:i+1]
# 存储到缓存
cache.store_prefill(layer_idx, sid, k_seq, v_seq)
# 返回最终的 hidden states(用于生成第一个 token 或直接输出)
return hidden_states
3.3 与 FlashAttention 结合的优化¶
在 prefill 阶段,长序列的计算瓶颈在于注意力矩阵的 O(N²) 显存。使用 FlashAttention 时,我们无需存储完整的注意力矩阵,但仍需要将 K、V 写出到缓存以供后续 decode 使用。因此,只需在 FlashAttention kernel 内部将计算好的 K、V 张量按块写回预先分配好的缓存区域即可。
# 伪代码:FlashAttention prefill 并写回缓存
for i in range(0, seq_len, block_size):
Qi = Q[:, :, i:i+block_size]
# 计算当前块与所有 K/V 的注意力(online softmax)
# ... 计算得到输出 Oi
# 同时,对于每个块 j,如果这是第一次计算,则将 Kj, Vj 写入缓存
if store_cache:
cache.k[layer, batch_idx, :, j_start:j_end] = Kj
cache.v[layer, batch_idx, :, j_start:j_end] = Vj
3.4 变长序列的处理¶
当 batch 内序列长度不同时,通常将输入左填充(或右填充)到相同长度,同时使用 key_padding_mask 屏蔽填充部分。在存储 KV 时,可以只存储有效长度部分,以减少缓存浪费。
def prefill_variable_length(model, input_ids, seq_ids, cache, padding_mask):
B, S = input_ids.shape
outputs, kv_list = model(input_ids)
for layer_idx, (k, v) in enumerate(kv_list):
for i, sid in enumerate(seq_ids):
valid_len = padding_mask[i].sum().item()
k_valid = k[i, :, :valid_len] # 截取有效部分
v_valid = v[i, :, :valid_len]
cache.store_prefill(layer_idx, sid, k_valid.unsqueeze(0), v_valid.unsqueeze(0))
return outputs
实现 decode 阶段单步推理的 KV Cache 更新:输入为当前 token 的 key/value,将其追加到对应序列的缓存尾部,并返回更新后的完整 KV 用于注意力计算。¶
4.1 Decode 阶段的特点¶
-
每次只输入一个 token(
input_ids形状为[B, 1])。 -
需要从缓存中读取该序列的全部历史 K、V,与当前的 Q 进行注意力计算。
-
计算完毕后,将当前 token 的 K、V 追加到缓存尾部。
4.2 代码实现(基于 PreallocKVCache)¶
def decode_step(model, token_ids, seq_ids, cache: PreallocKVCache):
"""
token_ids: [B, 1] 当前 token
返回每层的注意力输出
"""
B = token_ids.size(0)
# 1. 嵌入与位置编码(伪代码)
# hidden = embed(token_ids) + pos_embed
# 2. 逐层处理
for layer_idx in range(model.num_layers):
# 计算当前 token 的 Q, K, V
q = model.layers[layer_idx].q_proj(hidden) # [B, num_heads, 1, head_dim]
k = model.layers[layer_idx].k_proj(hidden)
v = model.layers[layer_idx].v_proj(hidden)
# 3. 从缓存中读取历史 KV
# 首先确保所有序列已注册
all_k, all_v = [], []
for i, sid in enumerate(seq_ids):
slot = cache.seq_to_slot[sid]
cur_len = cache.seq_len[slot]
# 读取该序列至此的全部 KV
k_hist = cache.k[layer_idx, slot, :, :cur_len] # [num_heads, cur_len, head_dim]
v_hist = cache.v[layer_idx, slot, :, :cur_len]
all_k.append(k_hist.unsqueeze(0))
all_v.append(v_hist.unsqueeze(0))
# 由于各序列长度可能不同,需要 padding 到相同长度或逐个计算
# 这里简化假设 batch 内所有序列长度相同,或者我们逐个序列计算注意力(效率低)
max_len = max(t.size(2) for t in all_k)
padded_k = torch.zeros(B, model.num_heads, max_len, model.head_dim, device=k.device)
padded_v = torch.zeros_like(padded_k)
mask = torch.zeros(B, max_len, dtype=torch.bool, device=k.device)
for i in range(B):
L = all_k[i].size(2)
padded_k[i, :, :L] = all_k[i]
padded_v[i, :, :L] = all_v[i]
mask[i, :L] = True
# 4. 注意力计算
scale = model.head_dim ** 0.5
scores = torch.matmul(q, padded_k.transpose(-2, -1)) / scale
scores = scores.masked_fill(~mask.unsqueeze(1).unsqueeze(2), float('-inf'))
attn_weights = torch.softmax(scores, dim=-1)
attn_output = torch.matmul(attn_weights, padded_v) # [B, H, 1, D]
# 5. 更新缓存:将当前 K、V 追加到各序列
for i, sid in enumerate(seq_ids):
cache.append_decode(layer_idx, sid, k[i:i+1], v[i:i+1])
# 6. 合并多头并经过 output projection
attn_output = attn_output.transpose(1,2).contiguous().view(B, 1, -1)
hidden = model.layers[layer_idx].out_proj(attn_output)
# 接着 FFN 等...
return hidden
4.3 优化点¶
-
避免逐序列 padding:当批量解码时,可以将所有序列的缓存放进一个连续张量,但序列长度不同会导致浪费。实践上常用连续批处理(continuous batching),将同一长度的序列组成一个 batch,减少 padding。
-
PagedAttention:通过分页机制,不同序列的 KV 块可以非连续存储,GPU kernel 通过块表直接索引,无需 padding。
设计一个支持动态序列长度的 KV Cache 容器,要求能够处理同一 batch 内不同序列长度的情况,并实现按 mask 读取有效位置。¶
5.1 设计思路¶
我们需要一个容器能够存储多个序列,每个序列长度独立。读取时可以返回一个 batch 级别的 K/V 和对应的 mask,用于后续注意力计算。
5.2 基于字典的变长实现¶
class VariableLengthKVCache:
def __init__(self, num_layers, num_heads, head_dim):
self.num_layers = num_layers
self.num_heads = num_heads
self.head_dim = head_dim
# 结构:self.k[layer][seq_id] = tensor [num_heads, seq_len, head_dim]
self.k = [{} for _ in range(num_layers)]
self.v = [{} for _ in range(num_layers)]
def update(self, layer_idx, seq_id, new_k, new_v):
"""new_k: [num_heads, 1, head_dim]"""
if seq_id not in self.k[layer_idx]:
self.k[layer_idx][seq_id] = new_k
self.v[layer_idx][seq_id] = new_v
else:
self.k[layer_idx][seq_id] = torch.cat([self.k[layer_idx][seq_id], new_k], dim=1)
self.v[layer_idx][seq_id] = torch.cat([self.v[layer_idx][seq_id], new_v], dim=1)
def get_batch(self, layer_idx, seq_ids):
"""返回 padding 后的 K, V 和 mask,形状 [B, heads, max_len, dim]"""
k_list = []
v_list = []
lengths = []
for sid in seq_ids:
k = self.k[layer_idx][sid]
v = self.v[layer_idx][sid]
k_list.append(k)
v_list.append(v)
lengths.append(k.size(1))
B = len(seq_ids)
max_len = max(lengths)
device = k_list[0].device
padded_k = torch.zeros(B, self.num_heads, max_len, self.head_dim, dtype=k_list[0].dtype, device=device)
padded_v = torch.zeros_like(padded_k)
mask = torch.zeros(B, max_len, dtype=torch.bool, device=device)
for i in range(B):
L = lengths[i]
padded_k[i, :, :L] = k_list[i]
padded_v[i, :, :L] = v_list[i]
mask[i, :L] = True
return padded_k, padded_v, mask
5.3 性能分析¶
-
每次 batch 读取都需要 padding,产生大量冗余计算和显存浪费,尤其是序列长度分布差异较大时。
-
改进:使用 Ragged Tensor 或 PagedAttention 的非连续访问,避免 padding。
5.4 与连续批处理的结合¶
在实际推理引擎中,调度器会尽量将长度相近的序列组成一个 batch,以最小化 padding 开销。同时,KV Cache 管理器还负责内存的分配与回收。
用原生 PyTorch 实现一个简化版的多头自注意力(MHA)前向推理,要求显式使用并更新外部传入的 KV Cache。¶
6.1 实现要点¶
-
需要传入一个 KV Cache 对象,并支持 prefill 和 decode 两种模式。
-
prefill 时,将整个序列的 K、V 存入缓存。
-
decode 时,从缓存读取历史 K、V,并追加新 token 的 K、V。
6.2 代码实现¶
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MHAWithCache(nn.Module):
def __init__(self, d_model, num_heads):
super().__init__()
self.d_model = d_model
self.num_heads = num_heads
self.head_dim = d_model // num_heads
self.q_proj = nn.Linear(d_model, d_model)
self.k_proj = nn.Linear(d_model, d_model)
self.v_proj = nn.Linear(d_model, d_model)
self.out_proj = nn.Linear(d_model, d_model)
def forward(self, x, cache: VariableLengthKVCache, layer_idx: int, seq_ids: list, is_prefill: bool):
"""
x: [B, S, d_model] 当前输入(prefill 时 S>1,decode 时 S=1)
cache: 外部缓存
layer_idx: 当前层索引
seq_ids: 长度为 B 的序列ID列表
"""
B, S, _ = x.shape
# 投影并分头
q = self.q_proj(x).view(B, S, self.num_heads, self.head_dim).transpose(1, 2) # [B, H, S, D]
k = self.k_proj(x).view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
if is_prefill:
# 将整个序列的 K/V 存入缓存
for i, sid in enumerate(seq_ids):
# 拆分 batch,逐个序列存入
# 注意:实际 prefill 可能批量写入,这里简化为逐序列
for t in range(S):
k_t = k[i:i+1, :, t:t+1, :] # [1, H, 1, D]
v_t = v[i:i+1, :, t:t+1, :]
cache.update(layer_idx, sid, k_t[0], v_t[0])
# prefill 模式下不需要从缓存读取,直接计算自注意力(带因果mask)
# 但为了统一,也可以从缓存读取(等同于刚写入的)
# 这里仍使用输入x自己计算注意力作为演示
scale = self.head_dim ** 0.5
causal_mask = torch.tril(torch.ones(S, S, device=x.device)).view(1, 1, S, S)
scores = torch.matmul(q, k.transpose(-2, -1)) / scale
scores = scores.masked_fill(causal_mask == 0, float('-inf'))
attn = torch.softmax(scores, dim=-1)
out = torch.matmul(attn, v)
else:
# decode 模式:S=1
# 1. 将当前 token 的 K/V 追加到缓存
for i, sid in enumerate(seq_ids):
cache.update(layer_idx, sid, k[i:i+1, :, 0:1, :][0], v[i:i+1, :, 0:1, :][0])
# 2. 从缓存读取历史 K/V
hist_k, hist_v, mask = cache.get_batch(layer_idx, seq_ids)
# 3. 注意力计算(Q 仅当前 token)
scale = self.head_dim ** 0.5
scores = torch.matmul(q, hist_k.transpose(-2, -1)) / scale
# 应用 mask(屏蔽 padding 位置)
scores = scores.masked_fill(~mask.unsqueeze(1).unsqueeze(2), float('-inf'))
attn = torch.softmax(scores, dim=-1)
out = torch.matmul(attn, hist_v)
# 合并多头并输出
out = out.transpose(1, 2).contiguous().view(B, S, self.d_model)
return self.out_proj(out)
6.3 注意事项¶
-
上述 prefill 模式中我们既写入了缓存又用输入自己计算了注意力,这是为了与 decode 统一接口。实际上 prefill 时可以直接从模型输出获取 KV,不必再单独存储一步(但为了后续 decode 需要存下来)。
-
在真实场景中,prefill 和 decode 通常分开为两个函数,以提高效率和清晰度。
实现 GQA(分组查询注意力)下的 KV Cache 存储与读取逻辑:给定 num_heads_q、num_heads_kv,写出如何从 KV Cache 中正确地广播 key/value 以匹配查询头数。¶
7.1 GQA 下的缓存存储¶
在 GQA 中,Key 和 Value 只计算 num_kv_heads 个头,而 Query 计算 num_q_heads 个头。缓存只需存储 num_kv_heads 个头的 K 和 V,形状为 [num_kv_heads, seq_len, head_dim]。
7.2 广播读取逻辑¶
当需要将缓存的 K/V 用于注意力计算时,由于 Q 头数量更多,我们需要将每个 KV 头重复 num_groups = num_q_heads // num_kv_heads 次,使其形状变为 [num_q_heads, seq_len, head_dim]。这可以通过 repeat_interleave 实现。
def broadcast_kv_for_gqa(k_cache, v_cache, num_q_heads, num_kv_heads):
"""
k_cache: [num_kv_heads, L, head_dim]
v_cache: [num_kv_heads, L, head_dim]
返回: [num_q_heads, L, head_dim] 的 K 和 V
"""
assert num_q_heads % num_kv_heads == 0
groups = num_q_heads // num_kv_heads
# 在头维度上重复
k_broadcast = k_cache.repeat_interleave(groups, dim=0)
v_broadcast = v_cache.repeat_interleave(groups, dim=0)
return k_broadcast, v_broadcast
7.3 与缓存容器的集成¶
在 VariableLengthKVCache 中,我们按 num_kv_heads 存储。当执行注意力计算时,取出缓存后调用 broadcast_kv_for_gqa 即可。注意广播只是创建一个视图,不会复制数据,因此内存开销可以忽略。
写出 PagedAttention 中的物理块结构定义,以及逻辑块表(Block Table)的数据结构,并说明一个 batch 中多条序列的块表如何组织。¶
8.1 物理块结构¶
物理块是一个固定大小(通常 16 或 32 个 token)的存储单元,包含该块内所有 token 在某一层的所有头的 K 和 V。在 GPU 显存中,可以连续排列。
class PhysicalBlock:
def __init__(self, block_size: int, num_heads: int, head_dim: int):
# 存储格式: [2, num_heads, block_size, head_dim],索引 0=K, 1=V
self.data = torch.zeros(2, num_heads, block_size, head_dim, dtype=torch.float16, device='cuda')
self.is_free = True
在 vLLM 等实现中,物理块通常按层管理,每层独立分配。但为了简化,可以在一个块中同时存储所有层的 KV(即形状 [num_layers, 2, num_heads, block_size, head_dim]),这样可以减少块表项数量,但灵活性下降。
8.2 逻辑块表(Block Table)¶
每个序列维护一个列表 block_table,记录该序列的逻辑块到物理块索引的映射。例如,一个序列长度为 50,block_size=16,则需要 4 个逻辑块。block_table = [phy_idx0, phy_idx1, phy_idx2, phy_idx3],其中第 4 个块只使用了前 2 个位置(内部碎片)。
8.3 Batch 内的块表组织¶
在实际推理时,我们需要将 batch 中所有序列的块表整合为一个数据结构,以便 GPU kernel 高效访问。常见的方式是构建一个二维数组 block_tables,形状 [B, max_blocks_per_seq],其中 max_blocks_per_seq 是支持的最大逻辑块数。对于每个序列,用物理块索引填充,未使用的位置填 -1。
batch_block_table = torch.full((B, max_blocks), -1, dtype=torch.int32)
for i, seq in enumerate(sequences):
for j, blk in enumerate(seq.block_table):
batch_block_table[i, j] = blk
GPU kernel 通过此表,以及每个序列的 context_len(已生成的 token 数),可以计算出逻辑块索引、物理块索引,进而访问数据。
实现一个简单的块分配器:支持从空闲块池中分配一个块,并返回块索引;当块释放时将其归还池中,要求用 Python 完成。¶
9.1 基本分配器¶
class BlockAllocator:
def __init__(self, total_blocks: int):
# 空闲块索引栈,pop 操作 O(1)
self.free_blocks = list(range(total_blocks))
self.total = total_blocks
def allocate(self) -> int:
if not self.free_blocks:
raise MemoryError(f"All {self.total} blocks are in use.")
return self.free_blocks.pop()
def free(self, block_index: int):
if block_index < 0 or block_index >= self.total:
raise ValueError("Invalid block index")
self.free_blocks.append(block_index)
def available(self) -> int:
return len(self.free_blocks)
9.2 线程安全版本¶
推理服务通常是多线程环境(如 vLLM 的调度器),需要加锁保护。
import threading
class ThreadSafeBlockAllocator(BlockAllocator):
def __init__(self, total_blocks: int):
super().__init__(total_blocks)
self.lock = threading.Lock()
def allocate(self) -> int:
with self.lock:
return super().allocate()
def free(self, block_index: int):
with self.lock:
super().free(block_index)
9.3 与序列绑定的内存管理¶
实际使用时,我们会记录每个序列占用了哪些物理块。当序列完成或换出时,遍历其块表,释放所有块。
class SequenceBlockManager:
def __init__(self, allocator: BlockAllocator, block_size: int):
self.allocator = allocator
self.block_size = block_size
self.seq_blocks = {} # seq_id -> list of physical block indices
def allocate_blocks_for_seq(self, seq_id: int, num_blocks: int):
if seq_id in self.seq_blocks:
self.free_seq(seq_id)
blocks = [self.allocator.allocate() for _ in range(num_blocks)]
self.seq_blocks[seq_id] = blocks
return blocks
def free_seq(self, seq_id: int):
for blk in self.seq_blocks.pop(seq_id, []):
self.allocator.free(blk)
9.4 扩展:支持“写时复制”和共享¶
当多个序列共享前缀(如 beam search 或并行采样)时,物理块可以被多个序列的块表引用。此时需要引用计数,当引用计数归零时才真正释放块。这可以通过将 free_blocks 改为引用计数数组来实现。
class RefCountBlockAllocator:
def __init__(self, total_blocks):
self.ref_counts = [0] * total_blocks
self.free_blocks = list(range(total_blocks))
def allocate(self) -> int:
blk = self.free_blocks.pop()
self.ref_counts[blk] = 1
return blk
def incref(self, block_index: int):
self.ref_counts[block_index] += 1
def decref(self, block_index: int):
self.ref_counts[block_index] -= 1
if self.ref_counts[block_index] == 0:
self.free_blocks.append(block_index)
以上设计构成了 vLLM 等现代推理引擎中 KV Cache 管理的基石。
基于块分配器,实现一个序列的 KV Cache 管理器:当序列长度增长时自动分配新的物理块,并在逻辑块表中记录映射。¶
原理¶
PagedAttention 将每个序列的 KV Cache 划分为固定大小的“物理块”(如 16 个 token)。一个序列的逻辑 KV 序列通过“块表”(Block Table)映射到一组物理块。当序列长度增长时,如果当前最后一个物理块已满,需要从全局空闲块池中分配新的物理块,并将其索引追加到块表中。
实现¶
class SequenceKVCache:
def __init__(self, allocator, block_size, num_heads, head_dim):
self.allocator = allocator # 全局块分配器
self.block_size = block_size # 每个块容纳的 token 数
self.num_heads = num_heads
self.head_dim = head_dim
self.block_table = [] # 逻辑块 -> 物理块索引
self.seq_len = 0 # 当前序列总长度(token数)
# 为简化,假设物理块存储由全局管理器统一维护,此处不持有块数据
def _ensure_capacity(self, required_tokens):
"""确保有足够的物理块容纳 required_tokens 个 token"""
needed_blocks = (required_tokens + self.block_size - 1) // self.block_size
while len(self.block_table) < needed_blocks:
blk = self.allocator.allocate()
self.block_table.append(blk)
def append_tokens(self, k, v):
"""
追加 K/V,形状 [num_heads, num_new_tokens, head_dim]。
该函数仅负责块分配与逻辑映射,实际写入需由外部配合物理块完成。
"""
num_new = k.shape[1]
new_total = self.seq_len + num_new
self._ensure_capacity(new_total)
# 实际写入逻辑应在外部根据块表和偏移写入物理存储,此处省略
self.seq_len = new_total
def get_block_table(self):
return self.block_table
def get_num_blocks(self):
return len(self.block_table)
关键点:_ensure_capacity 在每次追加前被调用,它依据所需的总 token 数计算所需块数,不足则分配。分配通过全局的 BlockAllocator 完成,它维护空闲物理块列表。实际物理块的 K/V 数据通常由全局大张量(如 kv_cache)统一存储,此处管理器只维护映射。
写出在 PagedAttention 下,给定逻辑块表和物理块存储,如何从物理块中按 token 位置顺序读取一段连续的 key/value 序列。¶
原理¶
给定序列的块表 block_table(列表),全局物理存储 kv_cache 形状为 [2, num_blocks, num_heads, block_size, head_dim],以及要读取的起止 token 位置 start 和 end,我们需要提取出形状为 [num_heads, end-start, head_dim] 的连续 K、V 张量。由于数据分散在不同物理块中,需逐块读取并拼接。
实现¶
def read_kv_range(block_table, kv_cache, start, end, block_size):
"""
block_table: list[int] 物理块索引列表
kv_cache: Tensor [2, num_blocks, num_heads, block_size, head_dim]
start, end: 读取区间 [start, end)
block_size: 每块 token 数
返回: (K, V) 每个形状 [num_heads, end-start, head_dim]
"""
num_heads = kv_cache.shape[2]
head_dim = kv_cache.shape[4]
out_k = torch.empty(num_heads, end - start, head_dim, dtype=kv_cache.dtype, device=kv_cache.device)
out_v = torch.empty_like(out_k)
out_idx = 0
pos = start
while pos < end:
logic_blk = pos // block_size # 逻辑块索引
offset = pos % block_size # 块内偏移
phy_blk = block_table[logic_blk] # 物理块索引
# 本块可读取的长度
can_read = min(block_size - offset, end - pos)
# 从物理块中复制
out_k[:, out_idx:out_idx+can_read, :] = kv_cache[0, phy_blk, :, offset:offset+can_read, :]
out_v[:, out_idx:out_idx+can_read, :] = kv_cache[1, phy_blk, :, offset:offset+can_read, :]
out_idx += can_read
pos += can_read
return out_k, out_v
说明:此函数产生了连续缓冲区,便于后续注意力计算(如矩阵乘法)。但在高效 GPU 内核中,为避免显式拷贝开销,通常直接在 kernel 中按块加载到共享内存,跳过中间缓冲区。
实现分块 KV Cache 下的单条序列追加新 token 操作:如果当前最后一个物理块未满,则直接填入;若已满,则从空闲池分配新块并链接到序列中。¶
原理¶
追加单个 token 的 K/V(形状 [num_heads, 1, head_dim])时,检查序列当前长度 seq_len:
-
若
seq_len % block_size == 0,说明最后一个逻辑块已满(或尚无块),需分配新物理块。 -
否则,直接写入最后一个块的
offset = seq_len % block_size位置。
物理块的存储假定为全局大张量 kv_cache(形状 [2, num_blocks, num_heads, block_size, head_dim]),我们需要一个全局函数完成写入。
实现¶
def append_token(block_table, seq_len, kv_cache, allocator, new_k, new_v, block_size):
"""
block_table: list[int] (会被原地修改)
seq_len: 当前序列长度(写入前)
kv_cache: 全局物理存储
allocator: BlockAllocator 实例
new_k, new_v: [num_heads, 1, head_dim]
返回: 更新后的 seq_len
"""
if seq_len % block_size == 0:
# 需要新块
blk = allocator.allocate()
block_table.append(blk)
logic_idx = seq_len // block_size
offset = seq_len % block_size
phy_idx = block_table[logic_idx]
# 写入
kv_cache[0, phy_idx, :, offset, :] = new_k[:, 0, :]
kv_cache[1, phy_idx, :, offset, :] = new_v[:, 0, :]
return seq_len + 1
扩展:若一次追加多个 token(如 prefill),可循环调用此函数,或批量处理整个块,提高效率。
写出 PagedAttention 内核中“根据块表将物理块内 KV 复制到连续临时缓冲区”的伪代码,并说明这样做的目的。¶
目的¶
物理块在显存中可能分散,直接用于矩阵乘法会产生大量非连续内存访问,严重影响带宽。在内核中,先将一个物理块的 K 和 V 加载到共享内存(连续缓冲区),然后在共享内存上执行注意力计算,可以最大化访问效率。
伪代码(CUDA 风格)¶
__global__ void paged_attention_kernel(
float* Q, // [num_heads, head_dim]
int* block_table, // 逻辑块->物理块映射
float* kv_cache, // [2][total_blocks][num_heads][block_size][head_dim]
float* output,
int seq_len,
int block_size,
int head_dim
) {
int head_idx = blockIdx.x;
// 共享内存作为连续缓冲区
__shared__ float K_tile[16][128]; // block_size x head_dim
__shared__ float V_tile[16][128];
float O[128] = {0};
float l = 0.0, m = -INFINITY;
int num_blocks = (seq_len + block_size - 1) / block_size;
for (int b = 0; b < num_blocks; ++b) {
int phy_idx = block_table[b];
int tokens = min(block_size, seq_len - b * block_size);
// 协同加载 K 块和 V 块到共享内存
for (int i = threadIdx.x; i < tokens * head_dim; i += blockDim.x) {
int t = i / head_dim;
int d = i % head_dim;
K_tile[t][d] = kv_cache[0 * total_blocks * num_heads * block_size * head_dim +
phy_idx * num_heads * block_size * head_dim +
head_idx * block_size * head_dim +
t * head_dim + d];
V_tile[t][d] = kv_cache[1 * ... + t * head_dim + d];
}
__syncthreads();
// 在共享内存中计算 Q * K_tile^T,然后 online softmax 更新 O
for (int t = 0; t < tokens; ++t) {
float dot = 0.0;
for (int d = 0; d < head_dim; ++d) dot += Q[head_idx][d] * K_tile[t][d];
// ... online softmax 逻辑
}
__syncthreads();
}
// 写回 O
}
说明:此伪代码展示了 FlashAttention 与 PagedAttention 的核心融合——每次循环通过块表获取物理块索引,从全局内存批量加载 K、V 到共享内存,然后在共享内存中计算注意力,避免了构造中间完整矩阵。
实现 prefill 后为一条序列构建其块表的函数:给定序列长度 S 和块大小 B,计算需要多少个块并分配,记录每个块的起止 token 位置。¶
实现¶
def build_block_table_for_prefill(seq_len, block_size, allocator):
"""
返回: block_table (list of physical block indices)
可选的起止位置信息可根据索引推算,无需显式存储。
"""
num_blocks = (seq_len + block_size - 1) // block_size
block_table = []
for i in range(num_blocks):
blk = allocator.allocate()
block_table.append(blk)
# 可记录该块的 [start, end) 但通常根据 i 和 block_size 可即时计算
return block_table
应用:prefill 阶段一次性分配所有需要的块,随后按顺序写入 K、V。每个逻辑块 i 的 token 范围为 [i*block_size, min((i+1)*block_size, seq_len))。
设计一个在推理服务中支持多序列并发请求的 KV Cache 块管理器,要求能够根据序列长度变化动态增减分配块,并处理序列结束后的块回收。¶
设计¶
管理器维护全局空闲块池和每个序列的块表,提供:
-
allocate_blocks(seq_id, num_blocks):为序列分配物理块。 -
ensure_capacity(seq_id, required_len):确保序列有足够块容纳长度。 -
free_seq(seq_id):回收序列所有块。 -
物理存储为全局大张量
kv_cache [2, total_blocks, num_heads, block_size, head_dim]。
代码框架¶
class GlobalKVCacheManager:
def __init__(self, total_blocks, block_size, num_heads, head_dim):
self.allocator = BlockAllocator(total_blocks)
self.block_size = block_size
self.kv_store = torch.zeros(2, total_blocks, num_heads, block_size, head_dim, dtype=torch.float16)
self.seq_blocks = {} # seq_id -> list[int] 物理块索引
self.seq_lengths = {} # seq_id -> int 当前长度
def allocate_blocks(self, seq_id, num_blocks):
if seq_id in self.seq_blocks:
self.free_seq(seq_id) # 先释放旧块
blocks = [self.allocator.allocate() for _ in range(num_blocks)]
self.seq_blocks[seq_id] = blocks
self.seq_lengths[seq_id] = 0
return blocks
def ensure_capacity(self, seq_id, required_len):
current_blocks = len(self.seq_blocks.get(seq_id, []))
needed_blocks = (required_len + self.block_size - 1) // self.block_size
if needed_blocks > current_blocks:
for _ in range(needed_blocks - current_blocks):
blk = self.allocator.allocate()
self.seq_blocks[seq_id].append(blk)
def free_seq(self, seq_id):
for blk in self.seq_blocks.pop(seq_id, []):
self.allocator.free(blk)
self.seq_lengths.pop(seq_id, None)
def write_token(self, seq_id, k, v):
# 实现追加一个 token,内部调用 ensure_capacity 和直接写入
pass
动态增减:扩容通过 ensure_capacity 自动进行;因为 decode 阶段长度只增不减,通常不需要减容。回收在序列完成或被驱逐时执行。
用 Python 实现一个简单的 LRU 缓存驱逐策略,当块池耗尽时选择最久未使用的序列释放其所有物理块,并供新序列使用。¶
原理¶
当空闲块耗尽而新请求需要块时,选择一个最近最少使用的序列(非活跃),将其所有物理块回收。维护每个序列的 last_access_time,使用有序字典或最小堆实现。
实现¶
import time
import threading
class LRUBlockManager:
def __init__(self, total_blocks, block_size):
self.allocator = BlockAllocator(total_blocks)
self.block_size = block_size
self.seq_blocks = {} # seq_id -> list[block_idx]
self.last_access = {} # seq_id -> timestamp
self.lock = threading.Lock()
def allocate_for_seq(self, seq_id, num_blocks):
with self.lock:
needed = num_blocks
# 先尝试直接分配
while needed > 0:
if self.allocator.available() > 0:
blk = self.allocator.allocate()
self.seq_blocks.setdefault(seq_id, []).append(blk)
needed -= 1
else:
# 驱逐最久未使用的序列
if not self.last_access:
raise MemoryError("No sequences to evict")
victim = min(self.last_access, key=self.last_access.get)
if victim == seq_id:
raise MemoryError("Cannot evict self")
# 释放 victim 的所有块
for blk in self.seq_blocks.pop(victim, []):
self.allocator.free(blk)
del self.last_access[victim]
self.last_access[seq_id] = time.time()
return self.seq_blocks[seq_id]
def touch(self, seq_id):
with self.lock:
self.last_access[seq_id] = time.time()
def free_seq(self, seq_id):
with self.lock:
for blk in self.seq_blocks.pop(seq_id, []):
self.allocator.free(blk)
self.last_access.pop(seq_id, None)
应用:vLLM 等系统在显存不足时,会将非活跃序列的 KV 块换出到 CPU 内存,LRU 是选择换出序列的常用策略。
写出在连续批处理调度下,如何为新加入的请求分配 KV Cache 块并初始化逻辑块表,同时对已完成的请求回收其所有物理块。¶
流程¶
连续批处理在每轮 decode 后动态调整 batch。新请求到来时,根据其 prompt 长度预分配块;完成的请求立即回收块。
代码片段¶
def add_new_request(manager, seq_id, prompt_len):
num_blocks = (prompt_len + manager.block_size - 1) // manager.block_size
blocks = manager.allocate_blocks(seq_id, num_blocks)
# 后续 prefill 写入 KV 到这些块
return blocks
def remove_finished(manager, seq_id):
manager.free_seq(seq_id)
实际调度:调度器维护 running 队列,每步迭代后检查生成 EOS 或达到 max_tokens 的序列,调用 remove_finished。新请求从 waiting 队列取出,先 prefill 分配块,然后转入 running。
实现 copy-on-write 机制在并行采样中的应用:当主序列 fork 出多个子序列时,子序列共享父序列的物理块,直到某个子序列需要写入新 token 时才复制块。¶
原理¶
多个子序列共享相同的 prompt 前缀,它们的块表最初指向相同的物理块。使用引用计数跟踪每个物理块被多少序列共享。当某个序列需要修改一个共享块(即写入新 token 到该块的空闲位置或后续需要修改)时,触发 COW:复制该物理块,更新该序列的块表指向新块,并减少原块的引用计数。
实现¶
class COWBlockManager:
def __init__(self, total_blocks, block_size):
self.allocator = BlockAllocator(total_blocks)
self.ref_counts = [0] * total_blocks # 每个物理块的引用计数
self.block_size = block_size
# 物理存储等省略
def fork_sequence(self, parent_seq):
"""基于父序列创建子序列,共享所有块"""
child = SequenceKVCache(self.allocator, self.block_size, ...)
child.block_table = parent_seq.block_table.copy()
child.seq_len = parent_seq.seq_len
# 增加所有共享块的引用计数
for blk in child.block_table:
self.ref_counts[blk] += 1
return child
def write_token(self, seq, k, v):
"""写入一个 token,自动处理 COW"""
if seq.seq_len % self.block_size == 0:
# 需要新块
blk = self.allocator.allocate()
self.ref_counts[blk] = 1
seq.block_table.append(blk)
else:
logic_idx = seq.seq_len // self.block_size
phy_idx = seq.block_table[logic_idx]
if self.ref_counts[phy_idx] > 1:
# 需要 COW:复制该块
new_blk = self.allocator.allocate()
self.ref_counts[new_blk] = 1
# 复制数据
copy_block(phy_idx, new_blk) # 实现略
# 更新引用
self.ref_counts[phy_idx] -= 1
seq.block_table[logic_idx] = new_blk
phy_idx = new_blk
# 写入 token 到 phy_idx 的适当偏移
# ...
seq.seq_len += 1
注意:通常 prompt 部分的块是只读的,不会触发 COW;只有子序列在生成新 token 时,如果最后一个逻辑块与父序列共享且该块仍有空位,写入该块会触发 COW。这样可以最大化共享前缀的收益。
用伪代码描述 FlashAttention 风格的块状注意力计算中,如何结合 PagedAttention 的块表来加载 KV 块,并在 SRAM 中迭代计算注意力输出。¶
已在第13题详细给出伪代码,此处补充说明融合流程:
-
外层循环遍历 Q 的块(通常 Q 一次性处理整个序列或按块,但 decode 时 Q 长度为 1,可直接全加载)。
-
内层循环遍历 KV 的逻辑块,通过块表获取物理块索引。
-
从全局显存将物理块的 K、V 加载到共享内存(SRAM)。
-
在共享内存中计算局部注意力分数,执行 online softmax 更新累积输出。
-
循环结束,归一化输出。
关键:物理块在共享内存中是连续排列的,消除了非连续访问带来的性能损失。
给定两个序列共享相同前缀,设计一个方法使它们的逻辑块表的前几个块指向相同的物理块,实现前缀缓存的 KV 复用,并写出相关的查找和引用计数逻辑。¶
设计¶
-
维护一个全局“前缀哈希表”或“Radix树”,将 token 序列映射到物理块列表。
-
当新序列的 prompt 到来时,先计算其前缀与已缓存前缀的重叠块数。
-
对于重叠部分,新序列的块表直接指向那些物理块,并增加引用计数;对于剩余部分,分配新块。
-
物理块需要引用计数,当某个序列结束或不再需要前缀时,减少引用计数;计数归零时回收块。
实现(简化版)¶
class PrefixCacheManager:
def __init__(self, total_blocks, block_size):
self.allocator = BlockAllocator(total_blocks)
self.block_size = block_size
self.ref_counts = [0] * total_blocks
# 前缀匹配表:可以使用 token 序列的哈希作为键,存储该前缀对应的块表
# 实际系统多采用 Radix 树,这里用简单字典示意
self.prefix_map = {} # token_tuple -> list of physical block indices
def match_prefix(self, prompt_tokens):
"""返回匹配的前缀长度(token数)和对应的块表"""
# 遍历可能的哈希,实际可逐块查找
# 这里仅为示意:假设 prompt_tokens 已经被分块
# 返回匹配的块表部分和匹配的 token 数
matched_blocks = []
matched_len = 0
# 逐块检查...
return matched_blocks, matched_len
def allocate_with_prefix(self, seq_id, prompt_tokens):
matched_blocks, matched_len = self.match_prefix(prompt_tokens)
block_table = matched_blocks.copy()
# 增加匹配块的引用计数
for blk in block_table:
self.ref_counts[blk] += 1
# 为剩余 token 分配新块
remaining = prompt_tokens[matched_len:]
new_blocks = self._allocate_for_tokens(remaining)
block_table.extend(new_blocks)
# 缓存新前缀(如果需要)
self._cache_prefix(prompt_tokens, block_table)
return block_table
def _allocate_for_tokens(self, tokens):
num_blocks = (len(tokens) + self.block_size - 1) // self.block_size
return [self.allocator.allocate() for _ in range(num_blocks)]
def release_blocks(self, block_table):
for blk in block_table:
self.ref_counts[blk] -= 1
if self.ref_counts[blk] == 0:
self.allocator.free(blk)
实际应用:vLLM 通过自动前缀缓存(Automatic Prefix Caching)实现了类似机制,利用 token 序列的哈希值快速匹配前缀物理块,极大减少了相同 system prompt 或共享前缀的 prefill 开销。