flashinfer.logits_processor.LogitsProcessor

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

LogitsProcessor 定义可以应用于 logits 或概率的高级转换。每个处理器都会自动合法化为低级 OpParameterizedOp,这些可以进行类型检查、验证和融合,以实现最佳性能。用户可以扩展此类来实现他们自己的处理器。

参数:

**params (Any) – 编译时处理器特定的参数。

示例

>>> import torch
>>> from flashinfer.logits_processor import LogitsPipe, TopK, Sample, TensorType
>>> torch.manual_seed(42)
>>>
>>> # Create a pipeline that legalizes to a fused op.
>>> pipe = LogitsPipe([
...     TopK(),         # Top-k filtering on logits
...     Sample()        # Sample from the filtered distribution
... ], input_type=TensorType.PROBS)  # assume the input is probabilities
>>>
>>> pipe
LogitsPipe([TopK -> Sample], ops=[ProbsTopKOp -> ProbsSampleOp], compiled_ops=[FusedProbsTopKSampleOp])

注意

子类必须实现 legalize() 方法,将高级处理器转换为一个或多个具有特定输入/输出类型的低级运算符

__init__(**params: Any)

初始化处理器。

参数:

**params (Any) – 编译时处理器特定的参数。

方法

__init__(**params)

初始化处理器。

legalize(input_type)

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