news 2026/9/7 3:54:14

vLLM 自定义 Logits Processor 实战:从批级接口到请求级适配器的完整指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
vLLM 自定义 Logits Processor 实战:从批级接口到请求级适配器的完整指南

vLLM 自定义 Logits Processor 实战:从批级接口到请求级适配器的完整指南

【免费下载链接】vllmA high-throughput and memory-efficient inference and serving engine for LLMs项目地址: https://gitcode.com/GitHub_Trending/vl/vllm

本篇指南基于 vLLM 仓库中的 examples/features/logits_processor/README.md 及其配套示例脚本,讲解如何在离线推理中注入自定义 logits processor,在采样前改写模型的输出分布。读完本文,你将掌握三种官方示例模式——批级(batch-level)processor、请求级(request-level)processor 的包装、以及需要访问引擎配置的构造器增强包装——并能理解其底层的批状态同步机制、参数校验链路与平台兼容限制。

Logits processor 的定位是在每个 decode step 的 forward 之后、采样之前,对 batch 的 logits 张量(形状为[batch_size, vocab_size])做任意修改,从而实现 token 屏蔽(token masking)、受限解码、自定义采样策略等受控生成行为。vLLM 出于效率考虑在批级处理 logits,因此如果你的 processor 天然只针对单个请求(例如依赖每请求自定义参数),就需要按仓库示例所示的方式做适配包装。

核心接口:LogitsProcessor 抽象基类与 BatchUpdate

自定义 processor 的实现依据是 vllm/v1/sample/logits_processor/interface.py 中定义的抽象基类LogitsProcessor(见 LogitsProcessor 定义),它有四个必须关注的方法和一个可选钩子:

方法说明
__init__(vllm_config, device, is_pin_memory)构造器。vLLM 的接口要求这三个参数必须存在,即使不使用也要保留
apply(logits)对整个 batch 的 logits 张量做修改,返回更新后的张量(允许就地修改)
update_state(batch_update)在每次 forward 之前、批组成发生变化时被调用,用于同步每请求状态
is_argmax_invariant()声明该 processor 是否影响贪心采样(temperature=0)下的 argmax 结果
validate_params(sampling_params)(类方法,可选)校验采样参数,非法时可抛出ValueError,引擎边界会将其转换为VLLMValidationError,在线服务中表现为 HTTP 400

批状态变化的载体是BatchUpdate冻结数据类(BatchUpdate 定义):

  • batch_size:当前 persistent batch 中的请求数;
  • removed:被移除请求的批索引序列;
  • added:新增请求的四元组(index, params, prompt_tok_ids, output_tok_ids)。注意output_tok_ids是该请求运行中输出 token 列表的引用,processor 通过它始终能看到最新的已生成 token;
  • moved:批内请求移动的三元组(index 1, index 2, directionality),方向性区分单向移动(UNIDIRECTIONAL)与双向交换(SWAP)。

BatchUpdate由 BatchUpdateBuilder 在调度过程中累积并生成;当批组成无变化时,update_state收到的参数为None

注册方式:类对象、FQCN 字符串与插件入口

通过离线推理入口LLMlogits_processors参数注入自定义 processor(该参数定义见 entrypoints/llm.py):

from vllm import LLM llm = LLM( model="facebook/opt-125m", logits_processors=[DummyLogitsProcessor], # 传入类,而非实例 )

加载逻辑集中在 build_logitsprocs。从源码结构看,最终生效的 processor 集合按以下顺序拼接:

  1. 内建 processorsMinTokensLogitsProcessorLogitBiasLogitsProcessorMinPLogitsProcessor(BUILTIN_LOGITS_PROCESSORS),分别支撑min_tokenslogit_biasmin_p等采样参数;
  2. 插件形式的 processors:通过 entry point 组vllm.logits_processors从已安装包动态加载(_load_logitsprocs_plugins),便于以独立分发的形式提供 processor;
  3. 用户显式指定的 processorslogits_processors参数可以是混合列表——既可以直接传入已加载的LogitsProcessor子类,也可以传完全限定类名字符串(FQCN),语法为<module>:<type>,例如"my_pkg.logitproc:MyProcessor",由 _load_logitsprocs_by_fqcns 负责导入并逐级定位类对象。

实例化时 vLLM 会对每个类调用ctor(vllm_config, device, is_pin_memory),并把实例按is_argmax_invariant()的返回值分桶到 LogitsProcessors 容器的argmax_invariant/non_argmax_invariant两个列表中——这一分桶意味着在贪心采样路径下,argmax 不变的 processor 可以被跳过,属于接口设计上的性能考量。

此外,每次提交请求时,引擎会遍历所有已加载的 processor 并调用其validate_params(sampling_params)做参数校验(validate_logits_processors_parameters)。

示例一:批级 processor(custom.py)

examples/features/logits_processor/custom.py 演示直接实现批级接口。示例中的DummyLogitsProcessor的行为是:当请求通过SamplingParams.extra_args传入target_token时,屏蔽除该 token 外的所有 logits,使每一步都只解码出目标 token。

python examples/features/logits_processor/custom.py

核心实现拆解:

from vllm import LLM, SamplingParams from vllm.config import VllmConfig from vllm.v1.sample.logits_processor import ( BatchUpdate, LogitsProcessor, ) from vllm.v1.sample.logits_processor.builtin import process_dict_updates class DummyLogitsProcessor(LogitsProcessor): """Fake logit processor to support unit testing and examples""" @classmethod def validate_params(cls, params: SamplingParams): target_token: Any | None = params.extra_args and params.extra_args.get( "target_token" ) if target_token is not None and not isinstance(target_token, int): raise ValueError( f"target_token value {target_token} {type(target_token)} is not int" ) def __init__( self, vllm_config: VllmConfig, device: torch.device, is_pin_memory: bool ): self.req_info: dict[int, int] = {} # 批索引 -> 目标 token id def is_argmax_invariant(self) -> bool: return False # 屏蔽 token 会改变贪心采样的 argmax 结果 def update_state(self, batch_update: BatchUpdate | None): def extract_extra_arg(params: SamplingParams) -> int | None: self.validate_params(params) return params.extra_args and params.extra_args.get("target_token") process_dict_updates( self.req_info, batch_update, # 根据请求细节计算 per-request 状态;返回 None 表示该 # processor 不适用于此请求 lambda params, _, __: extract_extra_arg(params), ) def apply(self, logits: torch.Tensor) -> torch.Tensor: if not self.req_info: return logits # 保存目标位置的原值 cols = torch.tensor( list(self.req_info.values()), dtype=torch.long, device=logits.device ) rows = torch.tensor( list(self.req_info.keys()), dtype=torch.long, device=logits.device ) values_to_keep = logits[rows, cols].clone() # 整行置 -inf,再恢复目标 token logits[rows] = float("-inf") logits[rows, cols] = values_to_keep return logits

几个实现要点值得注意:

  • process_dict_updates工具函数(builtin.py 实现)是批级 processor 的状态同步脚手架:它按照added → removed → moved的顺序维护一个dict[批索引, 状态]。你只需提供一个new_state回调——根据SamplingParams(以及可选的 prompt ids、输出 ids)返回该请求的状态,返回None即表示 processor 对该请求不生效。新增、移除、交换请求都会自动反映到字典中,这正是批级接口能"只影响部分请求"的关键。
  • apply的批量语义logits的第一维索引与 persistent batch 中的请求一一对应,因此通过rows/cols索引张量可以一次处理所有命中请求,避免逐请求循环。
  • 混合批构造:示例构造了 4 条 prompt,其中 50% 的请求携带target_token(取值 128 与 67),其余请求不带该参数。由于temperature=0.0,带参数的请求每步都输出同一个 token(如 token 67 在 OPT 词表中对应also),不带参数的请求则正常贪心解码。示例头部 docstring 中给出的预期输出即体现了这种对比。
prompts = [ "Hello, my name is", "The president of the United States is", "The capital of France is", "The future of AI is", ] sampling_params_list = [ SamplingParams(temperature=0.0, extra_args={"target_token": 128}), SamplingParams(temperature=0.0), SamplingParams(temperature=0.0, extra_args={"target_token": 67}), SamplingParams(temperature=0.0), ] llm = LLM(model="facebook/opt-125m", logits_processors=[DummyLogitsProcessor]) outputs = llm.generate(prompts, sampling_params_list)

示例二:包装请求级 processor(custom_req.py)

如果你的 processor 是按"单请求"粒度编写的——比如经典的f(output_ids, logits) -> logits签名——直接塞给批级接口并不合适。examples/features/logits_processor/custom_req.py 演示了如何用AdapterLogitsProcessor基类(AdapterLogitsProcessor 实现)把它适配为批级 processor:

python examples/features/logits_processor/custom_req.py
from vllm.v1.sample.logits_processor import ( AdapterLogitsProcessor, RequestLogitsProcessor, ) class DummyPerReqLogitsProcessor: """请求级 processor:屏蔽除 target_token 外的所有 logits""" def __init__(self, target_token: int) -> None: self.target_token = target_token def __call__( self, output_ids: list[int], logits: torch.Tensor, ) -> torch.Tensor: val_to_keep = logits[self.target_token].item() logits[:] = float("-inf") logits[self.target_token] = val_to_keep return logits class WrappedPerReqLogitsProcessor(AdapterLogitsProcessor): @classmethod def validate_params(cls, params: SamplingParams): target_token: Any | None = params.extra_args and params.extra_args.get( "target_token" ) if target_token is not None and not isinstance(target_token, int): raise ValueError(f"target_token value {target_token} is not int") def is_argmax_invariant(self) -> bool: return False def new_req_logits_processor( self, params: SamplingParams, ) -> RequestLogitsProcessor | None: target_token: Any | None = params.extra_args and params.extra_args.get( "target_token" ) if target_token is None: return None # 未提供 target_token 的请求不应用该 processor return DummyPerReqLogitsProcessor(target_token)

适配器的使用约定(源码 docstring 明确要求):

  1. 子类化AdapterLogitsProcessor
  2. 实现抽象方法new_req_logits_processor(params):根据该请求的SamplingParams返回一个定制化的请求级 processor 实例;返回None表示跳过该请求;
  3. 实现is_argmax_invariant()
  4. 一般不需要覆写__init__

从 AdapterLogitsProcessor 的实现 看,基类替你完成了全部批级簿记:

  • update_state内部仍走process_dict_updates,其new_state回调调用你的new_req_logits_processor,并把结果封装成functools.partial存进req_info(批索引 → partial 的映射);
  • partial 会预填充已生成 token 列表作为入参。这里有一个签名自适应细节:如果请求级 processor 的__call__接受 3 个参数(即f(prompt_ids, output_ids, logits)形式),基类会要求提供 prompt token ids 并一并传入;2 参形式则只传output_ids
  • apply时逐行取logits[req_idx]交给对应请求的 processor,若返回了新张量则回填到原行。由于 partial 持有输出 token 列表的引用,processor 每步看到的output_ids始终是该请求截至当前的完整生成序列。

示例三:需要引擎配置的请求级包装(custom_req_init.py)

examples/features/logits_processor/custom_req_init.py 覆盖一种特殊场景:请求级 processor 在初始化阶段就需要引擎配置或设备信息(例如按平台类型启用/禁用)。此时子类必须覆写包装基类的__init__(vllm_config, device, is_pin_memory),且覆写中应调用super().__init__(...)

python examples/features/logits_processor/custom_req_init.py
class WrappedPerReqLogitsProcessor(AdapterLogitsProcessor): """示例:覆写 __init__ 以获取设备类型信息""" @classmethod def validate_params(cls, params: SamplingParams): target_token = params.extra_args and params.extra_args.get("target_token") if target_token is not None and not isinstance(target_token, int): raise ValueError( f"`target_token` has to be an integer, got {target_token}." ) def __init__( self, vllm_config: VllmConfig, device: torch.device, is_pin_memory: bool ): super().__init__(vllm_config, device, is_pin_memory) self.is_cuda = device.type == "cuda" # 在构造期固化平台判断 def is_argmax_invariant(self) -> bool: return False def new_req_logits_processor( self, params: SamplingParams, ) -> RequestLogitsProcessor | None: if ( not self.is_cuda or ( target_token := params.extra_args and params.extra_args.get("target_token") ) is None ): return None return DummyPerReqLogitsProcessor(target_token)

示例用"非 CUDA 平台自动禁用 processor"建模了一个真实需求:is_argmax_invariant()new_req_logits_processor的决策逻辑依赖device。预期行为是——在 CUDA 设备上,带target_token的请求每步重复同一 token;而在非 CUDA 设备上,第 1、3 条请求会退化为正常贪心解码,因为 processor 对这些请求返回了None。除构造器之外,脚本的prompts/sampling_params_list/main()结构与示例二完全一致,可直接对照阅读。

关键概念对照与源码佐证

批级 vs 请求级的选择:vLLM 在 persistent batch 上以批级处理 logits,这是吞吐效率的前提。若你的 processor 逻辑天然按请求隔离(如每请求一个约束求解器、一个 token 过滤器),推荐走AdapterLogitsProcessor路线,把批簿记交给基类;只有当你需要在多个请求的 logits 行间做联合计算(跨行归一化、批量掩码等)时,才值得直接实现批级LogitsProcessor接口并利用索引张量做批量操作。

SamplingParams.extra_args传参约定:三个示例都通过extra_args={"target_token": ...}以请求粒度传递自定义参数。这是一个通用透传字典,vLLM 本身不消费其中的键,只负责原样带到 processor 侧;这也是为什么每个 processor 都要在validate_params中自行校验类型(示例中校验target_token必须是 int)。

DummyLogitsProcessor参考实现:按示例脚本的说明,DummyLogitsProcessor同时存在于一份测试参考实现中(vllm/test_utils.py),可以作为编写自定义 processor 的起点。本目录三个示例脚本内联了各自的简化版本,逻辑与其一致。

内建 processor 的对照参考:vLLM 内建的MinPLogitsProcessorLogitBiasLogitsProcessorMinTokensLogitsProcessor本身就展示了批级接口的最佳实践——例如 MinPLogitsProcessor 用 pinned CPU 张量 + 异步 H2D 拷贝批量同步每请求的min_p值,LogitBiasLogitsProcessor 用async_tensor_h2d把偏置压平为一维索引张量后一次性logits[rows, cols] += bias。阅读它们对写出高性能自定义 processor 有直接参考价值。

使用限制与兼容性边界

以下限制均可在 build_logitsprocs 源码 及 _load_custom_logitsprocs 中确认:

  • Pooling 模型不支持:对 embedding/pooling 类模型初始化logits_processors会直接抛出"Pooling models do not support custom logits processors."ValueError,且此时跳过全部 logits processor 加载;
  • 与推测解码互斥:启用 speculative decoding 时,自定义 logits processor 会触发ValueError(提示"Custom logits processors are not supported when speculative decoding is enabled."),并且min_plogit_bias参数在此模式下同样不生效(引擎仅保留MinTokensLogitsProcessor处理min_tokens);
  • TPU 平台暂不支持:当前 V1 TPU 路径下_load_custom_logitsprocs直接返回空列表,自定义 logits processor 不会被加载;
  • 参数校验异常语义validate_params中抛出的ValueError会在引擎边界被转换为VLLMValidationError,在线服务场景下对应 HTTP 400 响应,而不是进程级错误。

此外,logits_processors参数本身允许"类对象 + FQCN 字符串"混合列表(如[MyProcessor, "pkg.mod:AnotherProcessor"]),加载失败会带原始异常链抛出RuntimeError,便于定位导入错误。

小结:文件清单与延伸阅读

围绕本主题,仓库中值得深入阅读的文件:

文件作用
examples/features/logits_processor/README.md本指南对应的原始说明文档
examples/features/logits_processor/custom.py批级 processor 完整示例
examples/features/logits_processor/custom_req.py请求级 processor 的适配器包装示例
examples/features/logits_processor/custom_req_init.py构造期依赖引擎配置/设备的包装示例
vllm/v1/sample/logits_processor/interface.pyLogitsProcessor抽象基类与BatchUpdate定义
vllm/v1/sample/logits_processor/init.py加载链、build_logitsprocsAdapterLogitsProcessor
vllm/v1/sample/logits_processor/builtin.py内建 processor 与process_dict_updates工具
vllm/v1/sample/logits_processor/state.pyBatchUpdateBuilderLogitsProcessors容器
docs/features/logits_processors.md面向用户的 Logits Processor 特性文档

落地路径建议:先复制custom.py的批级骨架把apply/update_state跑通,再按需切换到custom_req.py的适配器模式降低状态管理成本;如果构造期需要引擎信息,则参照custom_req_init.py覆写__init__。编写前务必核对 pooling / 推测解码 / TPU 三类限制,并在validate_params中显式校验extra_args的类型与取值,这样自定义 processor 才能安全地进入 vLLM 的批采样流水线。

【免费下载链接】vllmA high-throughput and memory-efficient inference and serving engine for LLMs项目地址: https://gitcode.com/GitHub_Trending/vl/vllm

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/7 3:54:03

基于 Flask 和 GCN 的垃圾评论识别系统构建与部署

基于 Flask 和 GCN 的垃圾评论识别系统&#xff0c;是一个典型的“文本分类 Web 服务封装”项目。它解决的问题很明确&#xff1a;让用户通过 HTTP 接口提交评论文本&#xff0c;后端调用训练好的图卷积网络模型&#xff0c;判断这条评论是正常评论还是垃圾评论。适合正在做毕…

作者头像 李华
网站建设 2026/9/7 3:51:03

WIP服务端复活测试指南:GFDM XG2部署与接口验证实战

这次我们来看一个标记为 WIP 的服务器复活测试项目&#xff1a;GFDM XG2 服务器复活测试。所谓“复活测试”&#xff0c;通常指某个旧服务端因为依赖失效、配置丢失或代码停滞而无法运行&#xff0c;现在需要把它重新拉起来&#xff0c;验证核心流程是否还能走通。WIP 意味着代…

作者头像 李华
网站建设 2026/9/7 3:49:15

梅达焊接控制器实用指南:参数设定、故障排查与维护要点

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 3:47:43

JVM垃圾回收机制详解

JVM垃圾回收机制详解 1. 引言 1.1 什么是垃圾回收机制&#xff1f; 垃圾回收&#xff08;Garbage Collection&#xff0c;GC&#xff09;是JVM自动管理内存的一种机制。它负责回收不再使用的对象所占用的内存空间&#xff0c;避免内存泄漏&#xff0c;确保程序能够高效运行。 1…

作者头像 李华