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)。支持的数据类型:float32、float16、bfloat16。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])