flashinfer.logits_processor.LogitsProcessor¶
- class flashinfer.logits_processor.LogitsProcessor(**params: Any)¶
LogitsProcessor 定义可以应用于 logits 或概率的高级转换。每个处理器都会自动合法化为低级
Op或ParameterizedOp,这些可以进行类型检查、验证和融合,以实现最佳性能。用户可以扩展此类来实现他们自己的处理器。- 参数:
**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)将处理器合法化为低级运算符列表。