flashinfer.logits_processor.TopP

class flashinfer.logits_processor.TopP(**params: Any)

Top-p (nucleus) 过滤处理器。

保留累积概率达到阈值 p 的 token。

TensorType.PROBS -> TensorType.PROBS

参数:

top_p (floattorch.Tensor, Runtime) – 累积概率阈值,范围在 (0, 1] 内。可以是一个标量或每个批次的张量。

示例

>>> import torch
>>> from flashinfer.logits_processor import LogitsPipe, Softmax, TopP, Sample
>>> torch.manual_seed(42)
>>> pipe = LogitsPipe([TopP()])
>>> probs = torch.randn(2, 2, device="cuda")
>>> probs_normed = probs / probs.sum(dim=-1, keepdim=True)
>>> probs_normed
tensor([[ 0.0824,  0.9176], [-0.2541,  1.2541]], device='cuda:0')
>>> topp_probs = pipe(probs_normed, top_p=0.9)
>>> topp_probs
tensor([[0., 1.], [0., 1.]], device='cuda:0')
__init__(**params: Any)

TopP 处理器的构造函数。不需要编译时参数。

方法

__init__(**params)

TopP 处理器的构造函数。

legalize(input_type)

将处理器合法化为低级运算符列表。