资讯动态

FlashInfer logits_processor 声明式流水线指南:用 LogitsPipe 组合 LLM 采样算子与自定义融合规则

发布时间:2026/10/9 2:40:44 来源:尧图企业网站定制
大模型深度学习算子库后端高性能计算【免费下载链接】flashinferFlashInfer: Kernel Library for LLM Serving项目地址https://gitcode.com/gh_mirrors/fl/flashinfer点击查看免费下载FlashInfer 的logits_processor模块提供了一套声明式、可插拔的 LLM 输出处理流水线框架开发者只需用LogitsPipe把Temperature、Softmax、TopP、Sample等处理器像积木一样串起来框架就会自动完成类型检查、合法性校验与算子融合最终调用 FlashInfer 的高性能采样内核。读完本文你将掌握流水线的三段式编译原理类型推断 → 合法化 → 编译融合、所有内置处理器的参数语义以及如何通过子类化和自定义FusionRule扩展属于自己的处理器与融合策略。从 logits 到 token为什么要一套流水线框架LLM 解码decode阶段的核心流程是模型前向输出每个 token 的 logits未归一化的实数得分随后经过温度缩放、softmax 归一化、top-p/top-k 过滤等一系列变换最终采样得到整数 token ID。在 FlashInfer 中这一系列操作被抽象为一个可声明、可编译、可融合的流水线框架相关 API 文档位于 docs/api/logits_processor.rst其 Python 实现集中在 flashinfer/logits_processor/ 目录下包含 10 个模块文件pipeline.pyLogitsPipe流水线容器与执行逻辑processors.pyTemperature、Softmax、TopK、TopP、MinP、Sample等高阶处理器operators.py合法化后得到的低阶算子含融合算子types.pyTensorType语义类型与TaggedTensor张量包装compiler.py、fusion_rules.py、validators.py编译、融合规则与合法性检查legalization.py从处理器到算子的合法化阶段这一设计的核心理念是用户描述做什么框架负责怎么做——用户无需关心底层 CUDA kernel 的分派细节也无需手动拼接多次 kernel 调用。快速上手用 LogitsPipe 构建采样流水线LogitsPipe是流水线的核心容器。文档给出的最小示例展示了完整的构建与调用方式import torch from flashinfer.logits_processor import LogitsPipe, Temperature, Softmax, TopP, Sample # 创建流水线按顺序声明处理器 pipe LogitsPipe([ Temperature(), # 温度缩放 logits Softmax(), # logits 转概率 TopP(), # top-p核采样过滤 Sample() # 从分布中采样 ]) # 执行流水线 batch_size 4 vocab_size 5 logits torch.randn(batch_size, vocab_size, devicecuda) output_ids pipe(logits, temperature0.7, top_p0.9)注意这里的调用约定编译期参数在构建LogitsPipe时通过处理器构造函数传入运行时参数如temperature、top_p则在调用pipe(...)时以关键字参数传入。从 pipeline.py 的实现可以看出__call__会把这些运行时 kwargs 原样透传给流水线中的每一个算子LogitsPipe的输出永远是普通的torch.Tensortoken ID而不是带类型标记的TaggedTensor因此可以直接喂给下游的 PyTorch 代码。官方文档字符串中还有一个更贴近真实场景的示例见 pipeline.py使用 vocab_size32000 的 logits 张量torch.randn(4, 32000, devicecuda)配合TopK和Sample(deterministicTrue)构建流水线并传入temperature0.9, top_k40得到 4 个 token ID。从概率开始指定 input_type如果输入不是 logits 而是已经归一化的概率可以显式声明输入类型TensorType.PROBSfrom flashinfer.logits_processor import TensorType, TopK, Sample prob_pipe LogitsPipe( [TopK(), Sample()], input_typeTensorType.PROBS ) probs torch.softmax(logits, dim-1) token_ids prob_pipe(probs, top_k40)input_type参数很关键当第一个处理器既能接受 LOGITS 又能接受 PROBS例如TopK、Sample时框架无法自动推断输入类型必须显式指定否则会在流水线构建阶段抛出LegalizationError见 legalization.py 的infer_initial_type实现。三阶段流水线类型推断、合法化与编译LogitsPipe的构造函数pipeline.py内部按三步完成流水线的构建这也是理解整个框架的关键推断初始输入类型infer_initial_type尝试让第一个处理器分别对TensorType.LOGITS和TensorType.PROBS进行合法化根据其支持的类型集合确定流水线入口类型。合法化Legalizationlegalize_processors把声明式的高阶处理器逐个转换为带明确输入/输出类型声明的低阶算子Op。例如TopK()Sample()在 LOGITS 入口下会被转换为LogitsTopKOpLogitsSampleOp。编译Compilationcompile_pipeline对算子列表执行类型检查、合法性校验和融合优化得到compiled_ops。LogitsPipe构造函数的完整签名如下参数类型默认值说明processorsList[LogitsProcessor]必填按顺序应用的处理器列表不能为空compileboolTrue是否立即编译含融合优化为False时也可事后调用pipe.compile()input_typeOptional[TensorType]None期望的输入张量类型LOGITS 或 PROBS首处理器可同时接受两种类型时必须显式指定custom_fusion_rulesOptional[List[FusionRule]]None编译时额外应用的融合规则custom_validity_checksOptional[List[ValidityCheck]]None编译时额外应用的合法性检查编译阶段compiler.py依次执行类型检查_type_check遍历算子链确保前一个算子的OUT与下一个算子的IN严格匹配且首算子必须接受 LOGITS 或 PROBS 作为输入。合法性检查_run_validity_checks执行默认的校验规则详见下文。融合_fuse_all按规则优先级对算子窗口进行模式匹配与替换详见下文。若在未编译状态下直接调用pipe(...)框架会打出一条 warningPipeline is not compiled, running discrete ops.随后按未融合的离散算子列表逐一遍历执行pipeline.py。内置处理器详解参数、类型转换与语义所有内置处理器都定义在 processors.py它们继承自LogitsProcessor抽象基类并各自实现legalize(input_type)方法。各处理器的类型转换关系如下处理器输入类型输出类型编译期参数运行时参数TemperatureLOGITSLOGITS无temperature正浮点数或逐 batch 张量SoftmaxLOGITSPROBSenable_pdl可选默认 None 自动探测无TopKLOGITS 或 PROBSLOGITS 或 PROBSjoint_topk_topp默认 Falsetop_k正整数或逐 batch 张量TopPPROBSPROBS无top_p(0, 1] 内浮点数或逐 batch 张量MinPPROBSPROBS无min_p(0, 1] 内浮点数或逐 batch 张量SampleLOGITS 或 PROBSINDICESdeterministic默认 Trueindices、generator可选Temperature温度缩放Temperature的数学语义是logits / temperature即 logits 除以温度值。运行时参数temperature必须是正数可以是标量也可以是逐 batch 的张量用于 batch 内不同序列使用不同温度的场景。从算子实现 operators.py 可以看到非张量温度传入非正浮点数会直接抛出ValueError。Softmaxlogits 转概率Softmax将 LOGITS 转换为 PROBS官方文档特别强调一个流水线中最多只能出现一次 Softmax该约束由合法性检查规则 R1 强制保证见下文。其编译期参数enable_pdl用于控制底层 kernel 是否启用 PDLProgrammatic Dependent Launch默认值为None此时算子会通过device_support_pdl(tensor.data.device)自动探测当前设备是否支持 PDL 并据此启用见 operators.py。TopKtop-k 过滤TopK是少数既能处理 LOGITS 又能处理 PROBS 的处理器两种模式下的语义不同LOGITS 模式保留 top-k 得分最高的 token其余位置置为-inf对应flashinfer.sampling.top_k_mask_logits。PROBS 模式保留 top-k 概率最高的 token其余位置置 0 并重新归一化对应flashinfer.sampling.top_k_renorm_probs。编译期参数joint_topk_topp默认False决定当TopK后面紧跟TopP时是否启用top-k top-p 联合过滤的融合路径——该参数只有在融合规则执行时才会被读取见下文FusedProbsTopKTopPSampleOp。TopP核采样过滤TopPnucleus sampling只接受 PROBS 输入保留累积概率达到阈值top_p的 token 集合其余置 0 并重新归一化对应flashinfer.sampling.top_p_renorm_probs。top_p的取值范围是(0, 1]。MinPmin-p 过滤MinP同样只接受 PROBS 输入语义是保留概率不低于最大概率 ×min_p的 token——即过滤掉那些相对概率过低的 token。min_p必须位于(0, 1]可以是标量或逐 batch 张量。Sample采样生成 token IDSample是流水线的终点无论输入是 LOGITS 还是 PROBS输出都是TensorType.INDICES整数 token ID。由于 INDICES 是终止类型其后不允许再跟任何算子见合法性规则 R3。其编译期参数deterministic默认True选择确定性 kernel 实现运行时可选参数indices用于多个 batch 共享同一概率分布的批量采样场景generator用于传入torch.Generator以复现采样结果。底层分别对应flashinfer.sampling.sampling_from_logits与flashinfer.sampling.sampling_from_probs。类型系统TensorType 与 TaggedTensorTensorType流水线中的语义类型TensorType是一个枚举types.py定义了流水线中张量的三种语义LOGITS原始或掩码后的实数得分可为任意实数PROBS非负概率可能已归一化INDICES整数 token ID终止类型TaggedTensor携带类型信息的张量包装TaggedTensor是一个 frozen dataclasstypes.py内部持有data真正的torch.Tensor与type语义类型。它的作用是在算子链执行过程中维护语义类型信息以实现类型安全同时对外提供零摩擦的 PyTorch 互操作三个静态工厂方法TaggedTensor.logits(t)、TaggedTensor.probs(t)、TaggedTensor.indices(t)实现__torch_function__当对它调用任意 PyTorch 函数时会自动解包出底层data执行运算并返回普通torch.Tensor代理shape、device、dtype、size()等常用属性从使用角度看TaggedTensor主要供LogitsPipe内部使用用户传入普通张量后LogitsPipe.__call__会根据initial_type自动打上 LOGITS 或 PROBS 标签流水线内部逐算子传递TaggedTensor最终返回时解包为普通张量。每个算子Op见 op.py必须声明类属性IN与OUT例如IN TensorType.LOGITS、OUT TensorType.PROBS构造函数会强制校验这两个属性不能为空。ParameterizedOp进一步支持编译期默认参数调用时若运行时 kwargs 未提供某参数则回退到default_params中保存的默认值op.py。编译与融合从离散算子到高性能融合 kernel融合是这套框架性能收益的关键。Compiler._fuse_allcompiler.py在算子链上滑动窗口逐一比对内置及用户自定义融合规则的pattern一个算子类型元组匹配成功且guard通过后用build构造的融合算子替换窗口中的多个算子然后回退一步继续尝试直到无法再融合。默认融合规则内置规则定义在 fusion_rules.py模式pattern守卫guard融合产物优先级prio(TemperatureOp, SoftmaxOp)恒真FusedTemperatureSoftmaxOp100(ProbsTopKOp, TopPOp, ProbsSampleOp)joint_topk_topp TrueFusedProbsTopKTopPSampleOp100(ProbsTopKOp, ProbsSampleOp)恒真FusedProbsTopKSampleOp10(TopPOp, ProbsSampleOp)恒真FusedProbsTopPSampleOp10(MinPOp, ProbsSampleOp)恒真FusedProbsMinPSampleOp10规则按prio降序排列优先级越高越先尝试Compiler.register_fusion_rule中sort(keylambda r: -r.prio)。因此TemperatureSoftmax会融合为一个算子TopKSample、TopPSample、MinPSample分别融合而TopKTopPSample只有在TopK设置了joint_topk_toppTrue时才会走三算子联合融合路径。以 processors.py 的 doctest 为例流水线LogitsPipe([TopK(), Sample()], input_typeTensorType.PROBS)的__repr__会清晰展示三个阶段的结果LogitsPipe([TopK - Sample], ops[ProbsTopKOp - ProbsSampleOp], compiled_ops[FusedProbsTopKSampleOp])融合算子的底层 kernel 分派融合算子并非简单的依次调用而是直接进入 FlashInfer 的高性能采样路径。例如FusedTemperatureSoftmaxOp调用flashinfer.sampling.softmax(logits..., temperature..., enable_pdl...)在单一 kernel 内完成温度缩放与 softmaxoperators.py。源码注释特别说明融合后的流水线保持在公开分派路径上以确保与直接调用flashinfer.sampling.softmax的架构相关 softmax 路由一致。FusedProbsTopKSampleOp调用flashinfer.sampling.top_k_sampling_from_probs使用拒绝采样直接从 top-k 概率中采样operators.py避免了先过滤再采样两次 kernel 的开销。FusedProbsTopPSampleOp类似地调用flashinfer.sampling.top_p_sampling_from_probs。非融合的离散算子如ProbsTopKOp则调用top_k_renorm_probs等过滤 kernel并会为多 CTA kernel 分配一个 1MB 的row_states缓存缓冲区operators.py。合法性检查规则默认的合法性检查定义在 validators.pyR1 单 Softmax 规则single_softmax_rule流水线中SoftmaxOp出现次数不得超过 1 次否则编译失败。R3 INDICES 终止规则indices_terminal_rule任何输出 INDICES 的算子后面不得再跟算子。编译期发生的类型不匹配、规则违反会统一包装为ValueErrorPipeline creation failed: ... / Compilation failed: ...抛出便于在构建阶段尽早暴露错误。自定义处理器继承 LogitsProcessor 与 Op框架的扩展点之一是自定义处理器。文档给出的完整模式是处理器子类负责声明式描述legalize方法把处理器转译为带类型声明的算子from typing import Any, List from flashinfer.logits_processor import LogitsProcessor, Op, TensorType, TaggedTensor class CustomLogitsProcessor(LogitsProcessor): def __init__(self, **params: Any): super().__init__(**params) def legalize(self, input_type: TensorType) - List[Op]: return [CustomOp(**self.params)] class CustomOp(Op): # 定义输入输出张量类型 IN TensorType.LOGITS OUT TensorType.LOGITS def __call__(self, tensor: TaggedTensor, **kwargs: Any) - TaggedTensor: # 在这里实现实际运算逻辑 pass pipe LogitsPipe([CustomLogitsProcessor()]) # 流水线将被编译为 [CustomOp]要点回顾LogitsProcessor.__init__(**params)会把编译期参数保存在self.params中processors.py自定义处理器可沿用这一机制。子类必须实现抽象方法legalize(input_type)返回一个或多个OpOp必须声明IN/OUT类型并实现__call__。若你的Op需要运行时参数可继承ParameterizedOp并使用_get_param从 kwargs 或默认参数中取值。编译阶段会自动对你的算子执行类型检查——如果自定义Op的输出类型与下一个算子的IN不匹配编译会直接失败。自定义融合规则FusionRule 的三要素框架的另一个扩展点是自定义融合规则。FusionRule是一个 NamedTuplefusion_rules.py包含四个字段字段类型说明patternTuple[type, ...]要匹配的算子类型元组如(Temperature, Softmax)guardCallable[[List[Op]], bool]接收匹配到的算子窗口返回是否应用融合buildCallable[[List[Op]], Op]根据匹配到的算子构造融合算子可设置参数等prioint规则优先级越高越先尝试默认 0文档给出的完整示例def custom_fusion_guard(window: List[Op]) - bool: # 决定是否应用融合 return True def build_custom_fusion(window: List[Op]) - Op: # 构造融合算子例如设置参数等 return CustomOp() custom_rule FusionRule( pattern(Temperature, Softmax), guardcustom_fusion_guard, buildbuild_custom_fusion, prio20 ) pipe LogitsPipe( [Temperature(), Softmax(), Sample()], custom_fusion_rules[custom_rule] ) # 编译后的算子为 [CustomOp, Sample]注意示例中pattern(Temperature, Softmax)写的是处理器类型而编译阶段实际匹配的是合法化后的算子类型TemperatureOp、SoftmaxOp等。匹配通过isinstance完成compiler.py因此自定义融合规则中的 pattern 元素需要与你算子链中实际出现的算子类型对应。自定义规则通过LogitsPipe(custom_fusion_rules[...])或pipe.compile(custom_fusion_rules[...])注册与默认规则一同参与融合默认规则的优先级为 100 与 10自定义规则prio20会排在默认 10 之前、100 之后尝试。自定义合法性检查除了融合规则LogitsPipe还接受custom_validity_checks参数。ValidityCheck本质是Callable[[List[Op]], None]validators.py接收编译后的算子列表若不合规则抛出CompileError。你可以用它实现例如特定算子组合禁用依赖某种设备能力等业务级约束。测试验证编译路径与语义正确性框架的单元测试位于 tests/utils/test_logits_processor.py是理解框架行为边界的极佳参考编译路径一致性测试对比compileTrue与compileFalse两条路径的输出要求torch.allclose(..., atol1e-5)。参数组合覆盖 batch size ∈ {1, 99, 989}、vocab size ∈ {111, 32000, 128256}、多种分布正态、Gumbel与温度 ∈ {1.0, 0.5, 0.1}——这保证了融合后的 kernel 与离散算子组合在数值上等价。采样频率正确性对概率采样进行 500 万次试验统计采样频率与理论分布的一致性并验证被掩码为-inf的 token 永远不会被采样到。随机数可复现性通过克隆torch.Generator状态验证确定性 kernel 在相同随机种子下输出一致。实战建议与常见误区输入类型推断失败当流水线第一个处理器是TopK或Sample两者均可接受 LOGITS 与 PROBS时务必显式传入input_type否则构建会抛出LegalizationError。Softmax 只能出现一次一个流水线内只允许一个Softmax若输入已是概率直接以TensorType.PROBS入口构建省去 softmax 环节。Sample 必须放在最后Sample输出 INDICES 终止类型其后不能追加任何处理器。运行时参数按需透传temperature、top_k、top_p、min_p是运行时参数在pipe(...)调用时传入未提供且算子要求必填时会抛出ValueError。所有过滤参数都支持标量或逐 batch 张量便于实现 per-sequence 采样策略。融合收益默认规则会自动把TemperatureSoftmax、TopKSample、TopPSample、MinPSample融合为单 kernel 调用拒绝采样直接产出 token ID如需TopKTopPSample三算子联合融合记得设置TopK(joint_topk_toppTrue)。可复现采样采样可通过Sample(deterministicTrue)配合运行时generatortorch.Generator(...)获得可复现结果。整体来看logits_processor把 LLM 解码中最常见、最容易被反复摩擦的 logits 后处理环节收敛为一份声明即所得的流水线 DSL同时通过融合规则把性能关键路径交还给 FlashInfer 的采样 kernel 栈——这也是在 docs/api/logits_processor.rst 中反复强调的 declarative, pluggable framework 设计意图。赞分享大模型深度学习算子库后端高性能计算【免费下载链接】flashinferFlashInfer: Kernel Library for LLM Serving项目地址https://gitcode.com/gh_mirrors/fl/flashinfer点击查看免费下载相关推荐如何在5分钟内完成CCG多模型协作工作流的完整安装与配置如何在5分钟内完成CCG多模型协作工作流的完整安装与配置 你是否曾经在开发过程中需要频繁切换不同的AI模型来完成任务前端需要Gemini后端需要Codex人工智能AI 应用开发工具CLIAI Agentdsh-pluginDeepSeekTolaria Release Notes 阅读指南如何读懂日历版本号快速跟上每次更新Tolaria Release Notes 阅读指南如何读懂日历版本号快速跟上每次更新 Tolaria 是一款用于管理 Markdown 知识库的桌面应用桌面应用知识管理AI 应用MCP 服务MXNet Subgraph API 完全指南用 C 自定义子图搜索与算子融合MXNet Subgraph API 完全指南用 C 自定义子图搜索与算子融合 导读 Subgraph API 是 MXNet 官方提出的、用于将后端加速人工智能深度学习机器学习上一篇Sequoia未来路线图2025年Q4计划新增的5大策略与3项核心功能下一篇Minigrid入门指南5分钟快速掌握强化学习网格世界环境创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑