FlashInfer 中的 KV-Cache 布局¶
布局:NHD/HND¶
FlashInfer 为 KV-Cache 中的最后 3 个维度提供了两种布局:NHD 和 HND
NHD:最后 3 个维度组织为(seq_len, num_heads, head_dim)。HND:最后 3 个维度组织为(num_heads, seq_len, head_dim)。
由于 NHD 布局与 \(xW_k\) 和 \(xW_v\) 的输出一致,无需转置,因此该布局更自然。当 KV-Cache 使用低精度数据类型(例如 fp8)时,HND 布局更适合 GPU 实现。在实践中,我们没有观察到 fp16 kV-Cache 的这两种布局之间的显著性能差异,并且优先选择 NHD 布局以提高可读性。FlashInfer 在这两种布局上都实现了 Attention 内核,并提供了一个选项来在它们之间进行选择(默认情况下为 NHD)。
Ragged Tensor(不规则张量)¶
我们使用 Ragged Tensor 来存储 FlashInfer 中用于批量预填充自注意力机制的可变长度 Q/K/V 张量
在 Ragged Tensor 中,所有请求的 Q/K/V 都被打包到一个 data 张量中,没有填充。我们使用一个 indptr 数组(num_requests+1 个元素,第一个元素始终为零)来存储每个请求的可变序列长度信息(indptr[i+1]-indptr[i] 是请求 i 的序列长度),data 张量的形状为 (indptr[-1], num_heads, head_dim),当布局为 NHD 时。
我们可以使用 data[indptr[i]:indptr[i+1]] 来切片请求 i 的键(或值)。
注意
indptr 数组在 flashinfer 库中应为 int32 类型。 int64 类型的数组可能导致索引错误。
FlashInfer API¶
FlashInfer 提供了 flashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper 来计算存储在 Ragged Tensor 中的查询与存储在 Ragged KV-Cache 中的键/值之间的预填充注意力。
掩码布局(2D Ragged Tensor)¶
上述 Ragged Tensor 可以推广到多个“不规则”维度。例如,FlashInfer 中的注意力掩码对于批处理大小大于 1 时,是一个 2D 不规则张量
当请求数量大于 1 时,不同的请求可能具有不同的查询长度和 kv 长度。为了避免填充,我们使用 2D 不规则张量来存储注意力掩码。输入 qo_indptr 和 kv_indptr 数组(两者长度均为 num_requests+1)用于存储每个请求的可变序列长度信息,qo_indptr[i+1]-qo_indptr[i] 是请求 i 的查询长度(qo_len[i]),kv_indptr[i+1]-kv_indptr[i] 是请求 i 的 kv 长度(kv_len[i])。
所有请求的掩码数组被展平(查询作为第一维,kv 作为最后一维)并连接成一个 1D 数组:mask_data。FlashInfer 将隐式创建一个 mask_indptr 数组来存储展平掩码数组中每个请求掩码的起始偏移量:mask_indptr[1:] = cumsum(qo_len * kv_len)。
mask_data 的形状为 (mask_indptr[-1],),我们可以使用 mask_data[mask_indptr[i]:mask_indptr[i+1]] 来切片请求 i 的展平掩码。
为了节省内存,我们可以进一步将展平的布尔掩码数组打包成位打包数组(每个元素 1 位,8 个元素打包成一个 uint8),采用“小端”位序(有关更多详细信息,请参阅 numpy.packbits)。FlashInfer 接受布尔掩码和位打包掩码。如果提供布尔掩码,FlashInfer 将在内部将其打包成位打包数组。
FlashInfer API¶
flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper 和 flashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper 允许用户在 begin_forward 函数中指定 qo_indptr、kv_indptr 和自定义注意力掩码 custom_mask,掩码数据将在注意力内核中在 softmax(以及 softmax 缩放)之前添加到注意力分数中。
flashinfer.quantization.packbits() 和 flashinfer.quantization.segment_packbits() 是将布尔掩码打包成位打包数组的实用函数。
页表布局¶
当 KV-Cache 是动态的(例如,在追加或解码阶段),打包所有键/值效率不高,因为每个请求的序列长度会随着时间变化。vLLM 提出将 KV-Cache 组织为页表。在 FlashInfer 中,我们将页表视为一个块稀疏矩阵(每个使用的页可以看作是块稀疏矩阵中的一个非零块),并使用 CSR 格式 来索引 KV-Cache 中的页。
对于每个请求,我们记录其 page_indices、last_page_len,跟踪该请求使用的页和最后一页中的条目数。请求 i 的 KV 序列长度为 page_size * (len(page_indices[i]) - 1) + last_page_length[i]。
注意
每个请求的 last_page_len 必须大于零,并且小于或等于 page_size。
整体 kv_indptr 数组(长度为 num_requests+1)可以计算如下:[0, len(page_indices[0]), len(page_indices[0])+len(page_indices[1]), ...]。整体 kv_page_indices 数组(长度为 kv_indptr[-1])是所有请求的 page_indices 的连接。整体 kv_last_page_lens 数组(长度为 num_requests)是所有请求的 last_page_length 的连接。
kv_data 张量可以是单个 5D 张量,也可以是 4D 张量的元组。当存储在单个张量中时,kv_data 的形状为
kv_cache_nhd = torch.empty(max_num_pages, 2, page_size, num_heads, head_dim, dtype=torch.bfloat16) # NHD layout
kv_cache_hnd = torch.empty(max_num_pages, 2, num_heads, page_size, head_dim, dtype=torch.bfloat16) # HND layout
当存储在 4D 张量的元组中时,kv_data = (k_data, v_data),其中每个张量的形状为
k_cache_nhd = torch.empty(max_num_pages, page_size, num_heads, head_dim, dtype=torch.bfloat16) # NHD layout
k_cache_hnd = torch.empty(max_num_pages, num_heads, page_size, head_dim, dtype=torch.bfloat16) # HND layout
v_cache_nhd = torch.empty(max_num_pages, page_size, num_heads, head_dim, dtype=torch.bfloat16) # NHD layout
v_cache_hnd = torch.empty(max_num_pages, num_heads, page_size, head_dim, dtype=torch.bfloat16) # HND layout
其中 max_num_pages 是所有请求使用的最大页数,page_size 是每个页中适合的令牌数。单个张量存储中的 2 表示 K/V(第一个用于键,第二个用于值)。
注意
indptr 数组在 flashinfer 库中应为 int32 类型。 int64 类型的数组可能导致索引错误。 这也适用于 kv_page_indices 和 kv_last_page_lens 数组。
多头潜在注意力页布局¶
多头潜在注意力 (MLA) 是一种新的注意力机制,由 DeepSeek v2 提出,并用于后来的 DeepSeek 模型。MLA 将键缓存和值缓存统一到一个张量中,因此无需单独存储它们。与多头注意力或分组查询注意力相比,MLA 的 KV-Cache 没有 num_heads 维度,因此没有像 NHD 和 HND 布局这样的区别。
MLA 分离 RoPE(旋转位置编码)维度和其他头部维度。我们使用 kpe(带有位置编码的键)和 ckv(压缩的键/值)来命名这两个组件。用户可以将它们存储在单个 Paged KV-Cache 中
head_dim_ckv = 512
head_dim_kpe = 64
mla_paged_kv_cache = torch.empty(max_num_pages, page_size, head_dim_ckv + head_dim_kpe, dtype=torch.bfloat16)
ckv = mla_paged_kv_cache[:, :, :head_dim_ckv] # Slicing here does not copy or move data
kpe = mla_paged_kv_cache[:, :, head_dim_ckv:] # Slicing here does not copy or move data
并且 ckv 和 kpe 可以馈送到 MLA 注意力内核 flashinfer.mla.BatchMLAPagedAttentionWrapper。
FlashInfer API¶
flashinfer.page.append_paged_kv_cache() 可以将一批键/值(存储为不规则张量)追加到分页 KV-Cache 中(调用此 API 之前必须先为这些追加的键/值分配页面)。
flashinfer.decode.BatchDecodeWithPagedKVCacheWrapper 和 flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper 实现了存储在不规则张量中的查询与存储在分页 KV-Cache 中的键/值之间的解码注意力以及预填充/追加注意力。
多级级联推理数据布局¶
在使用多级 级联推理 时,查询和输出存储在不规则张量中,所有级别的 KV-Cache 存储在一个统一的分页 KV-Cache 中。每个级别都有一个唯一的 qo_indptr 数组,它是子树中要追加的累积令牌数量的前缀和,以及 kv_page_indptr、kv_page_indices 和 kv_last_page_len,其语义与 页面表布局 部分中的相同。下图介绍了如何为 8 个请求的追加注意力操作构建这些数据结构,我们将它们的 KV-Cache 视为用于前缀重用的 3 个级别
请注意,我们不必为每个级别的稀疏查询/输出张量或分页 kv-cache 更改数据布局。所有级别共享相同的底层数据布局,但我们使用不同的 qo_indptr / kv_page_indptr 数组,以便我们可以以不同的方式查看它们。
FlashInfer API¶
FlashInfer 提供了 flashinfer.cascade.MultiLevelCascadeAttentionWrapper 来计算级联注意力。
常见问题解答¶
- FlashInfer 如何管理 KV-Cache?
FlashInfer 本身不负责管理页面表(弹出和分配新页面等),我们将策略留给用户:不同的服务引擎可能具有不同的页面表管理策略。FlashInfer 仅负责计算查询与存储在 KV-Cache 中的键/值之间的注意力。