flashinfer.top_k¶
- flashinfer.top_k(input: Tensor, k: int, sorted: bool = False) Tuple[Tensor, Tensor]¶
基于基数的 Top-K 选择。
此函数从输入张量的每一行中选择最大的 k 个元素。它使用高效的基于基数的选择算法,该算法对于大型词汇表而言尤其快速。
此函数设计为
torch.topk的直接替代品,对于大型张量(词汇表大小 > 10000)具有更好的性能。- 参数:
input (torch.Tensor) – 输入张量,形状为
(batch_size, d),包含要选择的值。支持的数据类型:float32、float16、bfloat16。k (int) – 从每一行选择的 top 元素数量。
sorted (bool, optional) – 如果为 True,则返回的 top-k 元素将按降序排序。默认值为 False(未排序,速度更快)。
- 返回值:
values (torch.Tensor) – 形状为
(batch_size, k)的张量,包含 top-k 值。数据类型与输入相同。indices (torch.Tensor) – 形状为
(batch_size, k)且数据类型为 int64 的张量,包含 top-k 元素的索引。
注意
与
torch.topk不同,默认行为返回未排序的结果以提高性能。如果需要排序后的输出,请设置sorted=True。基于基数的算法在词汇表大小上为 O(n),与基于堆的方法的 O(n log k) 相比,使其对于大型词汇表更快。
对于小型词汇表(< 1000),
torch.topk可能更快。
示例
>>> import torch >>> import flashinfer >>> torch.manual_seed(42) >>> batch_size = 4 >>> vocab_size = 32000 >>> k = 256 >>> logits = torch.randn(batch_size, vocab_size, device="cuda") >>> values, indices = flashinfer.top_k(logits, k) >>> values.shape, indices.shape (torch.Size([4, 256]), torch.Size([4, 256]))
启用排序后(为了与 torch.topk 兼容)
>>> values_sorted, indices_sorted = flashinfer.top_k(logits, k, sorted=True) >>> # Values are now in descending order within each row
参见
torch.topkPyTorch 内置的 top-k 函数
sampling.top_k_mask_logitslogits 的 Top-k 掩码(将非 top-k 设置为 -inf)
sampling.top_k_renorm_probsTop-k 过滤和概率归一化