flashinfer.top_k_page_table_transform

flashinfer.top_k_page_table_transform(input: Tensor, src_page_table: Tensor, lengths: Tensor, k: int, row_to_batch: Tensor | None = None) Tensor

用于稀疏注意力机制的融合 Top-K 选择 + 页表转换。

此函数对输入分数执行 Top-K 选择,并在单个融合内核中通过页表查找转换选定的索引。用于稀疏注意力机制的第二阶段,其中选定的 KV 缓存位置需要通过页表映射。

对于每一行 i

output_page_table[i, j] = src_page_table[batch_idx, topk_indices[j]]

其中 batch_idx 由 row_to_batch[i] 确定(如果提供),否则为 i。

参数:
  • input (torch.Tensor) – 输入分数张量,形状为 (num_rows, max_len)。支持的数据类型:float32float16bfloat16

  • src_page_table (torch.Tensor) – 源页表,形状为 (batch_size, max_len),数据类型为 int32

  • lengths (torch.Tensor) – 每行的实际 KV 长度,形状为 (num_rows,),数据类型为 int32

  • k (int) – 从每一行选择的 top 元素数量。

  • row_to_batch (Optional[torch.Tensor], optional) – 从行索引到批索引的映射,形状为 (num_rows,),数据类型为 int32。如果为 None,则使用 1:1 映射(row_idx == batch_idx)。默认值为 None。

返回值:

output_page_table – 输出页表条目,形状为 (num_rows, k),数据类型为 int32。包含 Top-K 索引的收集到的页表条目。超出实际长度的位置设置为 -1。

返回值类型:

torch.Tensor

注意

  • 这专门设计用于稀疏注意力机制的第二阶段。

  • 如果 lengths[i] <= k,则输出简单地包含 src_page_table[batch_idx, 0:lengths[i]],其余位置设置为 -1。

示例

>>> import torch
>>> import flashinfer
>>> num_rows = 8
>>> max_len = 4096
>>> k = 256
>>> scores = torch.randn(num_rows, max_len, device="cuda", dtype=torch.float16)
>>> src_page_table = torch.randint(0, 1000, (num_rows, max_len), device="cuda", dtype=torch.int32)
>>> lengths = torch.full((num_rows,), max_len, device="cuda", dtype=torch.int32)
>>> output = flashinfer.top_k_page_table_transform(scores, src_page_table, lengths, k)
>>> output.shape
torch.Size([8, 256])