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),包含要选择的值。支持的数据类型:float32float16bfloat16

  • 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.topk

PyTorch 内置的 top-k 函数

sampling.top_k_mask_logits

logits 的 Top-k 掩码(将非 top-k 设置为 -inf)

sampling.top_k_renorm_probs

Top-k 过滤和概率归一化