FlashInfer 中的 KV-Cache 布局

布局:NHD/HND

FlashInfer 为 KV-Cache 中的最后 3 个维度提供了两种布局:NHDHND

  • 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 张量

Data structure of Ragged KV-Cache.

在 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 不规则张量

Data structure of Mask Layout.

当请求数量大于 1 时,不同的请求可能具有不同的查询长度和 kv 长度。为了避免填充,我们使用 2D 不规则张量来存储注意力掩码。输入 qo_indptrkv_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.BatchPrefillWithPagedKVCacheWrapperflashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper 允许用户在 begin_forward 函数中指定 qo_indptrkv_indptr 和自定义注意力掩码 custom_mask,掩码数据将在注意力内核中在 softmax(以及 softmax 缩放)之前添加到注意力分数中。

flashinfer.quantization.packbits()flashinfer.quantization.segment_packbits() 是将布尔掩码打包成位打包数组的实用函数。

页表布局

当 KV-Cache 是动态的(例如,在追加或解码阶段),打包所有键/值效率不高,因为每个请求的序列长度会随着时间变化。vLLM 提出将 KV-Cache 组织为页表。在 FlashInfer 中,我们将页表视为一个块稀疏矩阵(每个使用的页可以看作是块稀疏矩阵中的一个非零块),并使用 CSR 格式 来索引 KV-Cache 中的页。

Data structure of Paged KV-Cache.

对于每个请求,我们记录其 page_indiceslast_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_indiceskv_last_page_lens 数组。

多头潜在注意力页布局

多头潜在注意力 (MLA) 是一种新的注意力机制,由 DeepSeek v2 提出,并用于后来的 DeepSeek 模型。MLA 将键缓存和值缓存统一到一个张量中,因此无需单独存储它们。与多头注意力或分组查询注意力相比,MLA 的 KV-Cache 没有 num_heads 维度,因此没有像 NHDHND 布局这样的区别。

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

并且 ckvkpe 可以馈送到 MLA 注意力内核 flashinfer.mla.BatchMLAPagedAttentionWrapper

FlashInfer API

flashinfer.page.append_paged_kv_cache() 可以将一批键/值(存储为不规则张量)追加到分页 KV-Cache 中(调用此 API 之前必须先为这些追加的键/值分配页面)。

flashinfer.decode.BatchDecodeWithPagedKVCacheWrapperflashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper 实现了存储在不规则张量中的查询与存储在分页 KV-Cache 中的键/值之间的解码注意力以及预填充/追加注意力。

多级级联推理数据布局

在使用多级 级联推理 时,查询和输出存储在不规则张量中,所有级别的 KV-Cache 存储在一个统一的分页 KV-Cache 中。每个级别都有一个唯一的 qo_indptr 数组,它是子树中要追加的累积令牌数量的前缀和,以及 kv_page_indptrkv_page_indiceskv_last_page_len,其语义与 页面表布局 部分中的相同。下图介绍了如何为 8 个请求的追加注意力操作构建这些数据结构,我们将它们的 KV-Cache 视为用于前缀重用的 3 个级别

Cascade inference data layout.

请注意,我们不必为每个级别的稀疏查询/输出张量或分页 kv-cache 更改数据布局。所有级别共享相同的底层数据布局,但我们使用不同的 qo_indptr / kv_page_indptr 数组,以便我们可以以不同的方式查看它们。

FlashInfer API

FlashInfer 提供了 flashinfer.cascade.MultiLevelCascadeAttentionWrapper 来计算级联注意力。

常见问题解答

FlashInfer 如何管理 KV-Cache?

FlashInfer 本身不负责管理页面表(弹出和分配新页面等),我们将策略留给用户:不同的服务引擎可能具有不同的页面表管理策略。FlashInfer 仅负责计算查询与存储在 KV-Cache 中的键/值之间的注意力。