flashinfer.logits_processor.TopP¶
- class flashinfer.logits_processor.TopP(**params: Any)¶
Top-p (nucleus) 过滤处理器。
保留累积概率达到阈值 p 的 token。
TensorType.PROBS->TensorType.PROBS- 参数:
top_p (float 或 torch.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)将处理器合法化为低级运算符列表。