flashinfer.top_k_ragged_transform¶
- flashinfer.top_k_ragged_transform(input: Tensor, offsets: Tensor, lengths: Tensor, k: int) Tensor¶
用于稀疏注意力机制的融合 Top-K 选择 + 错位索引变换。
此函数在单个融合内核中执行输入分数的 Top-K 选择,并通过添加错位来变换选定的索引。用于稀疏注意力机制的第二阶段,具有错位/变长 KV 缓存。
- 对于每一行 i
output_indices[i, j] = topk_indices[j] + offsets[i]
- 参数:
input (torch.Tensor) – 输入分数张量,形状为
(num_rows, max_len)。支持的数据类型:float32,float16,bfloat16。offsets (torch.Tensor) – 每行添加的错位,形状为
(num_rows,),数据类型为int32。lengths (torch.Tensor) – 每行的实际 KV 长度,形状为
(num_rows,),数据类型为int32。k (int) – 从每一行选择的 top 元素数量。
- 返回值:
output_indices – 输出索引,形状为
(num_rows, k),数据类型为int32。包含 Top-K 索引加上错位。超出实际长度的位置设置为 -1。- 返回值类型:
torch.Tensor
注意
此函数专为稀疏注意力机制的第二阶段和错位 KV 缓存布局而设计。
如果 lengths[i] <= k,则输出包含 [offsets[i], offsets[i]+1, …, offsets[i]+lengths[i]-1],剩余位置设置为 -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) >>> offsets = torch.arange(0, num_rows * max_len, max_len, device="cuda", dtype=torch.int32) >>> lengths = torch.full((num_rows,), max_len, device="cuda", dtype=torch.int32) >>> output = flashinfer.top_k_ragged_transform(scores, offsets, lengths, k) >>> output.shape torch.Size([8, 256])