vLLM 自定义 Logits Processor 实战指南:编写、加载与批级请求级适配
2026/9/5 19:17:17 网站建设 项目流程

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

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

vLLM 的 logits processor(logits 处理器)允许你在采样前调整模型输出的下一 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_invariantnon_argmax_invariant两组,供采样路径按需分组调用。

另外可以指出一个实现层面的限制(源码证据,文档未展开):build_logitsprocs() 表明——

  • Pooling(池化/嵌入)模型不支持自定义 logits processor,初始化时直接报错;
  • 启用投机解码(speculative decoding)时拒绝加载自定义 logits processor,且min_plogit_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(f"target_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()), 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() # Mask all but target tokens logits[rows] = float('-inf') logits[rows, cols] = values_to_keep return logits

这个实现有两点值得学习:

  1. 稀疏(sparse)状态表示update_state()维护self.req_info字典,只记录指定了target_token的请求(键为批内索引,值为目标 token)。apply()中当req_info为空时直接原样返回 logits,实现整批短路;非空时则用向量化索引(logits[rows, cols])一次性完成掩码,而不是逐请求循环。
  2. 批状态同步update_state()依次处理 Add、Remove、Move 操作,使req_info的键始终与持久化批中的请求位置一致。

仓库中可直接运行的离线示例位于 examples/features/logits_processor/custom.py(python examples/features/logits_processor/custom.py),该目录的 README 还说明了custom_req.py(请求级包装器)与custom_req_init.py(需要引擎配置的请求级包装器)两个变体。

3.1 引擎如何构建BatchUpdate

原文档再次提示:该部分设计仍在演进,未来实现 logits processor 时可能不再需要考虑批状态变化。

update_state()的实现应假设模型运行器按如下模型更新持久化批状态(以BatchUpdate抽象表述):

  1. 识别当前引擎步中已完成的请求索引;
  2. 识别当前步新引入的请求;
  3. 用 Add 操作按"被替换请求索引从小到大"的顺序,尽量用新请求替换已完成的请求;
  4. 根据新请求与已完成请求数量的对比:
    1. 数量相同,进入下一步;
    2. 新请求更多:对未能替换已完成的剩余新请求执行 Add,索引从current_max_batch_index + 1起连续分配;
    3. 新请求更少
      • 对未被替换的已完成请求执行 Remove。这些被移除索引必然大于"上一步中被替换的已完成请求"的最大索引。Remove 可能使批进入非连续状态;
      • "压实(condense)"批为连续:从最低索引的空槽(由 Remove 造成)开始,执行一次单向 Move(Unidirectional Move),将批中当前最高非空槽的请求移入该空槽;随后按"空槽目标索引递增、非空槽来源索引递减"的顺序继续单向 Move,直到批恢复连续;
      • 收缩批:压实后,Remove 造成的空槽在批数组末尾聚成一块连续区域,因此将BatchUpdate.batch_size更新为非空槽数量。
  5. 为提升效率重排批:取决于注意力后端实现与批的当前特征,可能应用零次或多次 Swap Move 操作重排批。

关键约定(update_state()实现必须遵守):

  • 批更新操作必须按removes、adds、moves的顺序处理;
  • Add 操作的索引指Add 发生时刻的索引(即任何 Move 之前)。例如:某请求 Add 到索引 5 之后与索引 3 发生 swap,BatchUpdate.added中记录的仍是 5 而不是 3。换言之,Move 可被认为在 Add 和 Remove 之后应用;
  • Move 操作可按其在BatchUpdate.moved中出现的顺序假设依次应用;
  • 若没有新/已完成请求、也没有批重排,logits processor 收到的批更新就是None

对应源码印证:BatchUpdate 是一个 frozen dataclass,包含batch_sizeremoved(已完成请求索引序列)、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()addedmoved的调用,并在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 Processor(AdapterLogitsProcessor)

虽然 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( 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, ) -> 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 的唯一要求是:

  1. importlib.import_module()能解析点分路径部分并加载为模块;
  2. 类名部分能从已加载模块中导入;
  3. FQCN 指向的对象必须是LogitsProcessor的子类。

三个场景的写法:

# 1) LLM(离线) llm = LLM( model="facebook/opt-125m", logits_processors=["your.module.path:DummyLogitsProcessor"], )
# 2) AsyncLLM(异步引擎) engine_args = AsyncEngineArgs(model="facebook/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( model="facebook/opt-125m", logits_processors=[DummyLogitsProcessor], ) # 传给 AsyncLLM 构造器 engine_args = AsyncEngineArgs(model="facebook/opt-125m", logits_processors=[DummyLogitsProcessor]) async_llm = AsyncLLM.from_engine_args(engine_args)

7. 如何在请求中启用自定义 Logits Processor

是否需要按请求启用/禁用、以及需要传哪些参数,取决于处理器自身的设计。以DummyLogitsProcessor为例,用户通过自定义参数target_token来(1)为该请求启用处理器、(2)控制其行为:

REST API

curl 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_xargsextra_body传递):

batch = await client.completions.create( model="Qwen/Qwen2.5-1.5B-Instruct", ..., extra_body={ "vllm_xargs": { "target_token": 67 } } )

离线LLM

outputs_logitproc = llm.generate("your prompt", SamplingParams(..., extra_args={"target_token": 67}))

离线AsyncLLM

async for out in engine.generate(request_id="your request id", prompt="your prompt", sampling_params=SamplingParams(..., extra_args={"target_token": 67})): # Process async request outputs ...

8. 编写自定义 Logits Processor 的最佳实践

vLLM 初始化加载处理器后,每个引擎步都会对该处理器调用update_state()apply(),且两者都作用于持久化批中的所有请求——因此实现效率至关重要:

  • 在批粒度意识下写出高效的apply()update_state()

    • 尽量用向量化操作实现apply(),或在update_state()中批量更新内部状态向量;
    • 若处理器预计使用不频繁,适合采用"稀疏"表示:只保存"启用了处理器"的那些请求的元数据(如DummyLogitsProcessorreq_info字典);
    • 包装式请求级处理器无需自己实现这两个方法——AdapterLogitsProcessor的默认update_state()已维护稀疏状态(new_req_logits_processor()返回None的请求不进状态字典),默认apply()逐行顺序应用请求级处理器并组装输出张量。若默认实现性能不足,就放弃包装,改写为带向量化apply()/update_state()LogitsProcessor子类。
  • 由处理器作者决定的三件事

    1. 哪些按请求的属性配置处理器行为:你的update_state()重写决定了SamplingParams字段到处理器状态的映射;包装式处理器则由new_req_logits_processor()决定如何用SamplingParams初始化请求级实例。
    2. 按请求启用/禁用的条件:除非你的意图就是"对所有请求永远生效",否则应让处理器可被单请求禁用(例如把参数默认值设为None,或传入一个"什么都不做"的特定值如0.0),为禁用的请求省算力和内存;包装式处理器中,new_req_logits_processor()返回None即自动禁用该请求(基类默认实现已保证)。
    3. 批级短路的条件:即使支持按请求禁用,如果你用了对整批一次性运算的向量化实现,也很难因为"某一个请求禁用了"就省算(比如无法因单请求禁用就跳过apply()中整个向量化操作)。因此建议:设计apply()所有请求都禁用时直接返回未修改的输入张量;同理考虑update_state()在无请求启用时跳过步骤;一个简单的节省点是在batch_updateNone时提前返回。包装式处理器的基类默认已实现上述优化。
  • 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.py
BatchUpdateBuilder/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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询