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])