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)。支持的数据类型:float32float16bfloat16

  • 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_logitssoftmax 的组合应等效于 top_k_renorm_probs

参见

top_k_renorm_probs

top_k

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