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 字符串与插件入口
通过离线推理入口LLM的logits_processors参数注入自定义 processor(该参数定义见 entrypoints/llm.py):
from vllm import LLM llm = LLM( model="facebook/opt-125m", logits_processors=[DummyLogitsProcessor], # 传入类,而非实例 )加载逻辑集中在 build_logitsprocs。从源码结构看,最终生效的 processor 集合按以下顺序拼接:
- 内建 processors:
MinTokensLogitsProcessor、LogitBiasLogitsProcessor、MinPLogitsProcessor(BUILTIN_LOGITS_PROCESSORS),分别支撑min_tokens、logit_bias、min_p等采样参数; - 插件形式的 processors:通过 entry point 组
vllm.logits_processors从已安装包动态加载(_load_logitsprocs_plugins),便于以独立分发的形式提供 processor; - 用户显式指定的 processors:
logits_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.pyfrom 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 明确要求):
- 子类化
AdapterLogitsProcessor; - 实现抽象方法
new_req_logits_processor(params):根据该请求的SamplingParams返回一个定制化的请求级 processor 实例;返回None表示跳过该请求; - 实现
is_argmax_invariant(); - 一般不需要覆写
__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.pyclass 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 内建的MinPLogitsProcessor、LogitBiasLogitsProcessor、MinTokensLogitsProcessor本身就展示了批级接口的最佳实践——例如 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_p、logit_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.py | LogitsProcessor抽象基类与BatchUpdate定义 |
| vllm/v1/sample/logits_processor/init.py | 加载链、build_logitsprocs、AdapterLogitsProcessor |
| vllm/v1/sample/logits_processor/builtin.py | 内建 processor 与process_dict_updates工具 |
| vllm/v1/sample/logits_processor/state.py | BatchUpdateBuilder与LogitsProcessors容器 |
| 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),仅供参考