flashinfer.sampling.top_k_renorm_probs

flashinfer.sampling.top_k_renorm_probs(probs: Tensor, top_k: Tensor | int) Tensor

用于通过 top-k 阈值进行重整化的融合 GPU 内核。

参数:
  • probs (torch.Tensor) – 概率,形状 (batch_size, num_classes)。支持的数据类型:float32, float16, bfloat16

  • top_k (Union[torch.Tensor, int]) – 一个标量或形状为 (batch_size,) 的张量,表示用于重新归一化概率的 top-k 阈值,应在 (0, num_classes) 范围内。如果是一个标量,则对所有请求使用相同的阈值。如果是一个张量,则每个请求都有自己的阈值。我们保留 top-k 概率,将其余概率设置为零,并重新归一化概率。

返回值:

renorm_probs – 重新归一化的概率,形状 (batch_size, num_classes)。与输入 probs 具有相同的数据类型。

返回值类型:

torch.Tensor

示例

>>> import torch
>>> import flashinfer
>>> torch.manual_seed(42)
>>> batch_size = 4
>>> vocab_size = 5
>>> top_k = 3
>>> pre_norm_prob = torch.rand(batch_size, vocab_size).to(0)
>>> prob = pre_norm_prob / pre_norm_prob.sum(dim=-1, keepdim=True)
>>> prob
tensor([[0.2499, 0.2592, 0.1085, 0.2718, 0.1106],
        [0.2205, 0.0942, 0.2912, 0.3452, 0.0489],
        [0.2522, 0.1602, 0.2346, 0.1532, 0.2000],
        [0.1543, 0.3182, 0.2062, 0.0958, 0.2255]], device='cuda:0')
>>> renormed_probs = flashinfer.sampling.top_k_renorm_probs(prob, top_k)
>>> renormed_probs
tensor([[0.3201, 0.3319, 0.0000, 0.3480, 0.0000],
        [0.2573, 0.0000, 0.3398, 0.4028, 0.0000],
        [0.3672, 0.0000, 0.3416, 0.0000, 0.2912],
        [0.0000, 0.4243, 0.2750, 0.0000, 0.3007]], device='cuda:0')

注意

top_k_renorm_probssampling_from_probs 的组合应等效于 top_k_sampling_from_probs

参见

top_k_sampling_from_probs, sampling_from_probs, top_p_renorm_probs

top_k

通用的 top-k 选择(返回索引和值)