flashinfer.sampling.top_k_mask_logits¶
- flashinfer.sampling.top_k_mask_logits(logits: Tensor, top_k: Tensor | int) Tensor¶
用于通过 top-k 阈值进行掩码的融合 GPU 内核。
- 参数:
logits (torch.Tensor) – softmax 之前的 logits,形状为
(batch_size, num_classes)。支持的数据类型:float32、float16、bfloat16。top_k (Union[torch.Tensor, int]) – 要么是一个标量,要么是一个形状为
(batch_size,)的张量,表示用于掩码 logits 的 top-k 阈值,应在(0, num_classes)范围内。如果是一个标量,则对所有请求使用相同的阈值。如果是一个张量,则每个请求都有自己的阈值。我们保留 top-k logits,将其余值设置为负无穷大。
- 返回值:
masked_logits – 掩码后的 logits,形状为
(batch_size, num_classes)。与输入logits具有相同的数据类型。- 返回值类型:
torch.Tensor
示例
>>> import torch >>> import flashinfer >>> torch.manual_seed(42) >>> batch_size = 4 >>> vocab_size = 5 >>> top_k = 3 >>> logits = torch.randn(batch_size, vocab_size).to(0) >>> logits tensor([[ 1.9269, 1.4873, 0.9007, -2.1055, -0.7581], [ 1.0783, 0.8008, 1.6806, 0.3559, -0.6866], [-0.4934, 0.2415, -0.2316, 0.0418, -0.2516], [ 0.8599, -0.3097, -0.3957, 0.8034, -0.6216]], device='cuda:0') >>> masked_logits = flashinfer.sampling.top_k_mask_logits(logits, top_k) >>> masked_logits tensor([[ 1.9269, 1.4873, 0.9007, -inf, -inf], [ 1.0783, 0.8008, 1.6806, -inf, -inf], [ -inf, 0.2415, -0.2316, 0.0418, -inf], [ 0.8599, -0.3097, -inf, 0.8034, -inf]], device='cuda:0')
注意
top_k_mask_logits和softmax的组合应等效于top_k_renorm_probs。