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_probs和sampling_from_probs的组合应等效于top_k_sampling_from_probs。