vLLM 自定义 Logits Processor 实战指南:编写、加载与批级请求级适配

发布时间:2026/9/5 19:17:29
vLLM 自定义 Logits Processor 实战指南:编写、加载与批级请求级适配 vLLM 自定义 Logits Processor 实战指南编写、加载与批级请求级适配【免费下载链接】vllmA high-throughput and memory-efficient inference and serving engine for LLMs项目地址: https://gitcode.com/GitHub_Trending/vl/vllmvLLM 的 logits processorlogits 处理器允许你在采样前调整模型输出的下一 token 概率分布实现 token 掩码、约束解码、自定义采样策略等可控生成行为。本文基于 vLLM 官方文档 custom_logitsprocs 与仓库源码完整讲解如何编写一个自定义 logits processor、通过哪三种方式将其加载进 vLLM 引擎、如何为单个请求传参启用它以及官方给出的最佳实践读完后可直接实现一个可运行的批级batch-level或请求级request-level自定义处理器。注意原文档声明logits processor 的设计仍在演进中相关 API 近期可能发生变化vLLM 团队计划尽快稳定这部分 API。1. 背景Logits Processor 在 vLLM 中的工作粒度Logits processor 的作用是调整下一 token 的概率分布通常用于引导模型朝期望的行为方向发展。在 vLLM 中logits processor 工作在批batch粒度在某个引擎步engine step中logits processor 消费的是模型输出的原始 logits 张量形状为(num_requests) x (vocab_size)。对于所有启用了该处理器的请求处理器对 logits 张量中对应的行施加变换其余行保持不动变换后的 logits 张量随后被送入 softmax。这与 vLLM 0 时代的请求级设计要求处理器是一个Callable不同批级接口能显著提升性能。源码中批级抽象位于 LogitsProcessor 基类引擎侧的处理器集合封装在 LogitsProcessors 类中——它会在初始化时按is_argmax_invariant()的返回值把所有处理器分成argmax_invariant与non_argmax_invariant两组供采样路径按需分组调用。另外可以指出一个实现层面的限制源码证据文档未展开build_logitsprocs() 表明——Pooling池化/嵌入模型不支持自定义 logits processor初始化时直接报错启用投机解码speculative decoding时拒绝加载自定义 logits processor且min_p、logit_bias参数在投机解码下不生效从源码结构看当前 V1 引擎在TPU 平台上尚未支持自定义 logits processor_load_custom_logitsprocs()对 TPU 直接返回空列表。2. 编写自定义 Logits Processor必须实现的接口自定义 logits processor 必须继承vllm.v1.sample.logits_processor.LogitsProcessor抽象基类定义于 interface.py至少实现以下方法方法职责validate_params(cls, sampling_params)类方法。校验请求的SamplingParams尤其是自定义参数是否合法非法时抛ValueError。请求发送到入口点时会被调用非法请求会被直接拒绝。必须实现否则非法参数会导致处理器行为异常__init__(self, vllm_config, device, is_pin_memory)构造函数。vllm_config是引擎配置结构device是硬件加速器设备信息is_pin_memory指示 pinned memory 是否可用于辅助实现apply(self, logits) - torch.Tensor消费(num_requests) x (vocab_size)的 logits 张量在批粒度上施加变换并返回变换后的张量。可以原地in-place或离地out-of-place修改原地修改更省内存is_argmax_invariant(self) - bool若处理器永不改变某个请求中 logit 值最高的 token ID返回True否则返回False。该方法在启动时求值一次若返回True当某一步的所有请求都使用贪心采样时vLLM 会跳过对该处理器的调用update_state(self, batch_update)在引擎步开始时消费BatchUpdate数据结构据此更新处理器内部维护的持久化批状态。batch_update可能为None表示批成员无变化——此时仍可能需要根据此前保留的output_token_ids引用更新内部状态源码中validate_params的默认实现直接返回None即默认不校验这与文档强调务必自行实现相呼应classmethod def validate_params(cls, sampling_params: SamplingParams): Validate sampling params for this logits processor. Raise VLLMValidationError (preferred) / ValueError (backward compatible) for invalid params. return None请求侧校验入口为 validate_logits_processors_parameters()旧式处理器抛出的ValueError会在引擎边界被转换为VLLMValidationError从而让在线服务返回 HTTP 400。3. 完整示例DummyLogitsProcessor下面这个示例实现了一个简单的自定义处理器它消费(num_requests) x (vocab_size)的 logits 张量用float(-inf)掩掉除target_token之外的所有 token对未指定target_token的请求处理器处于禁用状态。它通过检查每个请求SamplingParams.extra_args中的target_token自定义参数来决定是否启用、以及保留哪个 token。import torch from vllm.config import VllmConfig from vllm.sampling_params import SamplingParams from vllm.v1.sample.logits_processor import (BatchUpdate, LogitsProcessor, MoveDirectionality) class DummyLogitsProcessor(LogitsProcessor): Fake logit processor to support unit testing and examples classmethod def validate_params(cls, params: SamplingParams): target_token: int | 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(ftarget_token value {target_token} is not int) def __init__(self, vllm_config: VllmConfig, device: torch.device, is_pin_memory: bool): self.req_info: dict[int, int] {} def is_argmax_invariant(self) - bool: Never impacts greedy sampling return False def update_state(self, batch_update: BatchUpdate | None): if not batch_update: return # Process added requests. for index, params, _, _ in batch_update.added: assert params is not None self.validate_params(params) if params.extra_args and (target_token : params.extra_args.get(target_token)): self.req_info[index] target_token else: self.req_info.pop(index, None) if self.req_info: # Process removed requests. for index in batch_update.removed: self.req_info.pop(index, None) # Process moved requests, unidirectional move (a-b) and swap # (a-b) for adx, bdx, direct in batch_update.moved: a_val self.req_info.pop(adx, None) b_val self.req_info.pop(bdx, None) if a_val is not None: self.req_info[bdx] a_val if direct MoveDirectionality.SWAP and b_val is not None: self.req_info[adx] b_val def apply(self, logits: torch.Tensor) - torch.Tensor: if not self.req_info: return logits # Save target values before modification cols torch.tensor( list(self.req_info.values()), dtypetorch.long, devicelogits.device ) rows torch.tensor( list(self.req_info.keys()), dtypetorch.long, devicelogits.device ) values_to_keep logits[rows, cols].clone() # Mask all but target tokens logits[rows] float(-inf) logits[rows, cols] values_to_keep return logits这个实现有两点值得学习稀疏sparse状态表示update_state()维护self.req_info字典只记录指定了target_token的请求键为批内索引值为目标 token。apply()中当req_info为空时直接原样返回 logits实现整批短路非空时则用向量化索引logits[rows, cols]一次性完成掩码而不是逐请求循环。批状态同步update_state()依次处理 Add、Remove、Move 操作使req_info的键始终与持久化批中的请求位置一致。仓库中可直接运行的离线示例位于 examples/features/logits_processor/custom.pypython examples/features/logits_processor/custom.py该目录的 README 还说明了custom_req.py请求级包装器与custom_req_init.py需要引擎配置的请求级包装器两个变体。3.1 引擎如何构建BatchUpdate原文档再次提示该部分设计仍在演进未来实现 logits processor 时可能不再需要考虑批状态变化。update_state()的实现应假设模型运行器按如下模型更新持久化批状态以BatchUpdate抽象表述识别当前引擎步中已完成的请求索引识别当前步新引入的请求用 Add 操作按被替换请求索引从小到大的顺序尽量用新请求替换已完成的请求根据新请求与已完成请求数量的对比数量相同进入下一步新请求更多对未能替换已完成的剩余新请求执行 Add索引从current_max_batch_index 1起连续分配新请求更少对未被替换的已完成请求执行 Remove。这些被移除索引必然大于上一步中被替换的已完成请求的最大索引。Remove 可能使批进入非连续状态压实condense批为连续从最低索引的空槽由 Remove 造成开始执行一次单向 MoveUnidirectional Move将批中当前最高非空槽的请求移入该空槽随后按空槽目标索引递增、非空槽来源索引递减的顺序继续单向 Move直到批恢复连续收缩批压实后Remove 造成的空槽在批数组末尾聚成一块连续区域因此将BatchUpdate.batch_size更新为非空槽数量。为提升效率重排批取决于注意力后端实现与批的当前特征可能应用零次或多次 Swap Move 操作重排批。关键约定update_state()实现必须遵守批更新操作必须按removes、adds、moves的顺序处理Add 操作的索引指Add 发生时刻的索引即任何 Move 之前。例如某请求 Add 到索引 5 之后与索引 3 发生 swapBatchUpdate.added中记录的仍是 5 而不是 3。换言之Move 可被认为在 Add 和 Remove 之后应用Move 操作可按其在BatchUpdate.moved中出现的顺序假设依次应用若没有新/已完成请求、也没有批重排logits processor 收到的批更新就是None。对应源码印证BatchUpdate 是一个 frozen dataclass包含batch_size、removed已完成请求索引序列、added(index, params, prompt_tok_ids, output_tok_ids)四元组序列与moved(index1, index2, directionality)三元组序列注释中明确写出操作应按 removed、added、moved 的顺序处理其中added里的output_tok_ids是对请求运行中输出 token 列表的引用使处理器始终能看到最新的已生成 token。构建侧则由 BatchUpdateBuilder 在每一步聚合removed_append()、added、moved的调用并在get_and_reset()中生成BatchUpdate无任何变化时返回None。4. 通过自定义参数Custom Arguments为处理器传参与内建处理器不同自定义处理器往往需要SamplingParams或 REST API 中并不存在的配置项。vLLM 的 自定义参数custom arguments机制正是为此设计自定义参数以字典形式传递增删都不需要重编译 vLLM。当然你也可以直接复用SamplingParams的既有字段视设计而定。离线写入SamplingParams.extra_args任何能访问SamplingParams的代码包括你的处理器都可见SamplingParams(extra_args{your_custom_arg_name: 67})在线OpenAI 兼容 REST API 与 Anthropic 兼容/v1/messages端点通过vllm_xargs传递底层会被赋值到SamplingParams.extra_args因此基于extra_args的处理器实现天然兼容离线/在线两种场景curl http://localhost:8000/v1/completions \ -H Content-Type: application/json \ -d { model: Qwen/Qwen2.5-1.5B-Instruct, ... vllm_xargs: {your_custom_arg: 67} }OpenAI SDK 用户则通过extra_body传递vllm_xargs。原文档提醒务必为你的自定义参数实现validate_params否则非法自定义参数可能引发不可预期的行为。5. 包装已有的请求级 Logits ProcessorAdapterLogitsProcessor虽然 vLLM 引擎以批粒度应用处理器但你可能想沿用为 vLLM v0 开发的请求级处理器——那种要求Callable形式类型定义见 vllm/logits_process.py、符合如下注解的实现RequestLogitsProcessor Union[ # (output token ids, logits tensor) - logits tensor Callable[[list[int], Tensor], Tensor], # (prompt token ids, output token ids, logits tensor) - logits tensor Callable[[list[int], list[int], Tensor], Tensor], ]请求级处理器在 vLLM 引擎中不被直接支持但 vLLM 提供了便捷的包装流程子类化AdapterLogitsProcessor即可把一个请求级Callable包装成兼容的批级处理器。包装时需要重写validate_params(cls, params)校验请求采样参数重写is_argmax_invariant(self)如实反映请求级处理器是否可能改变最高 logit token重写new_req_logits_processor(self, params)从SamplingParams创建新的请求级处理器实例返回None表示该请求不应用处理器。下例中DummyPerReqLogitsProcessor是你的请求级处理器的替身from vllm.v1.sample.logits_processor import ( AdapterLogitsProcessor, # Wrapper base-class RequestLogitsProcessor, # Request-level logitsproc type annotation ) # Stand-in for your request-level logits processor: class DummyPerReqLogitsProcessor: The request-level logits processor masks out all logits except the token id identified by target_token def __init__(self, target_token: int) - None: Specify target_token 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 # Example of wrapping the request-level logits processor: class WrappedPerReqLogitsProcessor(AdapterLogitsProcessor): Example of wrapping a fake request-level logit processor to create a batch-level logits processor 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( ftarget_token value {target_token} is not int ) def is_argmax_invariant(self) - bool: return False def new_req_logits_processor( self, params: SamplingParams, ) - Optional[RequestLogitsProcessor]: 返回针对该请求定制的请求级处理器 当请求未提供整数 target_token 时返回 None不应用 target_token: Any | None params.extra_args and params.extra_args.get( target_token ) if target_token is None: return None return DummyPerReqLogitsProcessor(target_token)AdapterLogitsProcessor基类源码帮你完成了大部分脏活这正是文档中包装类无需自己实现apply()/update_state()的来源基类用self.req_info: dict[int, partial[torch.Tensor]]维护稀疏状态new_req_logits_processor()返回None的请求不出现在字典里partial持有对output_ids列表的引用因此始终基于最新的已生成 token 运行默认update_state()复用process_dict_updates()同步 Add/Remove/Move并自动丢弃已完成请求的状态默认apply()按req_info中的请求索引逐行调用请求级处理器for req_idx, req_lp in self.req_info.items()若请求级处理器返回了新张量则就地写回对应行。如果这个默认的逐行循环性能不足文档建议不要包装而是把它重写为直接继承LogitsProcessor的批级实现用向量化方式实现apply()/update_state()。6. 三种加载自定义 Logits Processor 的方式处理器在初始化时加载。关键限制引擎加载完成后已加载的处理器集合不可变更也不能按请求动态加载新处理器。以下三种方式均适用于这一点约束之下。方式 1初始化时传入全限定类名FQCN该方式同时支持离线与在线场景。FQCN 格式为dotted.path.to.module:ClassName可以作为logits_processors参数传给LLM/AsyncLLM的 Python 构造器作为 CLI 参数传给vllm serve。vllm serve ... --logits_processors logits processor 1 logits processor 2 ...FQCN 的唯一要求是importlib.import_module()能解析点分路径部分并加载为模块类名部分能从已加载模块中导入FQCN 指向的对象必须是LogitsProcessor的子类。三个场景的写法# 1) LLM离线 llm LLM( modelfacebook/opt-125m, logits_processors[your.module.path:DummyLogitsProcessor], )# 2) AsyncLLM异步引擎 engine_args AsyncEngineArgs(modelfacebook/opt-125m, logits_processors[your.module.path:DummyLogitsProcessor]) async_llm AsyncLLM.from_engine_args(engine_args)# 3) vllm serve在线服务 vllm serve facebook/opt-125m --logits_processors your.module.path:DummyLogitsProcessor源码印证_load_logitsprocs_by_fqcns() 会把已加载的类 FQCN 字符串的混合列表统一转换为类列表对字符串按module_path:qualname拆分后importlib.import_module加载再沿点分路径getattr走到目标对象并断言其为LogitsProcessor子类任何加载失败都会抛出带上下文的RuntimeError。CLI 侧--logits-processors参数在 engine/arg_utils.py 中注册dest 为logits_processors与文档中logits_processors的 Python 参数名一致。方式 2以 Python Entry Point 自动发现通过 setuptools 的 entry points已安装的包可以把自己暴露为插件。vLLM 初始化时会自动扫描vllm.logits_processors这一 entry point 分组并加载其中所有处理器。若你的自定义处理器在一个 Python 包里只需在该包的pyproject.toml中为每个处理器声明一个 entry point[project.entry-points.vllm.logits_processors] dummy_logits_processor your.module.path:DummyLogitsProcessor包安装后每次 vLLM 初始化都会自动加载这些处理器无需再向LLM/AsyncLLM构造器或vllm serve显式传参。注意vLLM总是加载vllm.logits_processors分组下暴露的所有处理器不能选择性地只加载其中一部分。源码印证分组常量LOGITSPROCS_GROUP vllm.logits_processors与加载逻辑 _load_logitsprocs_plugins() 均位于 vllm/v1/sample/logits_processor/__init__.py单个 entry point 加载失败会抛出RuntimeError并记录错误日志。方式 3仅离线直接向构造器传 Python 类对象可以向LLM/AsyncLLM构造器传入一个或多个处理器类对象。这种方式最灵活类既可以在实例化LLM/AsyncLLM的同一源文件内本地定义也可以从 Python 包导入。# 从模块导入 from some.module import DummyLogitsProcessor # ...或者本地定义... from vllm.v1.sample.logits_processor import LogitsProcessor class DummyLogitsProcessor(LogitsProcessor): # 见上文 DummyLogitsProcessor 实现 ... # 传给 LLM 构造器 llm LLM( modelfacebook/opt-125m, logits_processors[DummyLogitsProcessor], ) # 传给 AsyncLLM 构造器 engine_args AsyncEngineArgs(modelfacebook/opt-125m, logits_processors[DummyLogitsProcessor]) async_llm AsyncLLM.from_engine_args(engine_args)7. 如何在请求中启用自定义 Logits Processor是否需要按请求启用/禁用、以及需要传哪些参数取决于处理器自身的设计。以DummyLogitsProcessor为例用户通过自定义参数target_token来1为该请求启用处理器、2控制其行为REST APIcurl http://localhost:8000/v1/completions \ -H Content-Type: application/json \ -d { model: Qwen/Qwen2.5-1.5B-Instruct, ... vllm_xargs: {target_token: 67} }OpenAI SDKvllm_xargs经extra_body传递batch await client.completions.create( modelQwen/Qwen2.5-1.5B-Instruct, ..., extra_body{ vllm_xargs: { target_token: 67 } } )离线LLMoutputs_logitproc llm.generate(your prompt, SamplingParams(..., extra_args{target_token: 67}))离线AsyncLLMasync for out in engine.generate(request_idyour request id, promptyour prompt, sampling_paramsSamplingParams(..., extra_args{target_token: 67})): # Process async request outputs ...8. 编写自定义 Logits Processor 的最佳实践vLLM 初始化加载处理器后每个引擎步都会对该处理器调用update_state()和apply()且两者都作用于持久化批中的所有请求——因此实现效率至关重要在批粒度意识下写出高效的apply()与update_state()尽量用向量化操作实现apply()或在update_state()中批量更新内部状态向量若处理器预计使用不频繁适合采用稀疏表示只保存启用了处理器的那些请求的元数据如DummyLogitsProcessor的req_info字典包装式请求级处理器无需自己实现这两个方法——AdapterLogitsProcessor的默认update_state()已维护稀疏状态new_req_logits_processor()返回None的请求不进状态字典默认apply()逐行顺序应用请求级处理器并组装输出张量。若默认实现性能不足就放弃包装改写为带向量化apply()/update_state()的LogitsProcessor子类。由处理器作者决定的三件事哪些按请求的属性配置处理器行为你的update_state()重写决定了SamplingParams字段到处理器状态的映射包装式处理器则由new_req_logits_processor()决定如何用SamplingParams初始化请求级实例。按请求启用/禁用的条件除非你的意图就是对所有请求永远生效否则应让处理器可被单请求禁用例如把参数默认值设为None或传入一个什么都不做的特定值如0.0为禁用的请求省算力和内存包装式处理器中new_req_logits_processor()返回None即自动禁用该请求基类默认实现已保证。批级短路的条件即使支持按请求禁用如果你用了对整批一次性运算的向量化实现也很难因为某一个请求禁用了就省算比如无法因单请求禁用就跳过apply()中整个向量化操作。因此建议设计apply()在所有请求都禁用时直接返回未修改的输入张量同理考虑update_state()在无请求启用时跳过步骤一个简单的节省点是在batch_update为None时提前返回。包装式处理器的基类默认已实现上述优化。update_state必须丢弃已完成请求的信息被 Add 替换或遭遇 Remove 的请求包装式处理器由基类默认处理。is_argmax_invariant()的用法若处理器行为恒定可硬编码返回True/False若不变性随用户配置动态变化也可程序化判定——正因如此该方法不是类方法而是实例方法启动时对每个实例求值一次见 LogitsProcessors 的两分组逻辑。9. 相关代码与文档索引内容路径本文对应的官方文档docs/features/custom_logitsprocs.md自定义参数custom arguments文档docs/features/custom_arguments.md处理器基类 /BatchUpdate/MoveDirectionality定义vllm/v1/sample/logits_processor/interface.py内建处理器MinTokens / LogitBias / MinP与AdapterLogitsProcessor、加载入口build_logitsprocsvllm/v1/sample/logits_processor/init.pyBatchUpdateBuilder/LogitsProcessors状态管理vllm/v1/sample/logits_processor/state.py内建处理器实现vllm/v1/sample/logits_processor/builtin.py请求级处理器类型注解v0 兼容vllm/logits_process.py--logits-processorsCLI 参数注册vllm/engine/arg_utils.py可运行离线示例批级 / 请求级 / 请求级引擎配置examples/features/logits_processor/按以上步骤你可以在不修改、不重编译 vLLM 源码的前提下为 vLLM 添加任意自定义的采样前 logits 变换逻辑并在离线与在线两种部署形态下按请求粒度启用它。【免费下载链接】vllmA high-throughput and memory-efficient inference and serving engine for LLMs项目地址: https://gitcode.com/GitHub_Trending/vl/vllm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考