ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

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-levelprocessor、请求级request-levelprocessor 的包装、以及需要访问引擎配置的构造器增强包装——并能理解其底层的批状态同步机制、参数校验链路与平台兼容限制。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 是否影响贪心采样temperature0下的 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 通过它始终能看到最新的已生成 tokenmoved批内请求移动的三元组(index 1, index 2, directionality)方向性区分单向移动UNIDIRECTIONAL与双向交换SWAP。BatchUpdate由 BatchUpdateBuilder 在调度过程中累积并生成当批组成无变化时update_state收到的参数为None。注册方式类对象、FQCN 字符串与插件入口通过离线推理入口LLM的logits_processors参数注入自定义 processor该参数定义见 entrypoints/llm.pyfrom vllm import LLM llm LLM( modelfacebook/opt-125m, logits_processors[DummyLogitsProcessor], # 传入类而非实例 )加载逻辑集中在 build_logitsprocs。从源码结构看最终生效的 processor 集合按以下顺序拼接内建 processorsMinTokensLogitsProcessor、LogitBiasLogitsProcessor、MinPLogitsProcessorBUILTIN_LOGITS_PROCESSORS分别支撑min_tokens、logit_bias、min_p等采样参数插件形式的 processors通过 entry point 组vllm.logits_processors从已安装包动态加载_load_logitsprocs_plugins便于以独立分发的形式提供 processor用户显式指定的 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。示例一批级 processorcustom.pyexamples/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( ftarget_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()), dtypetorch.long, devicelogits.device ) rows torch.tensor( list(self.req_info.keys()), dtypetorch.long, devicelogits.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其余请求不带该参数。由于temperature0.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(temperature0.0, extra_args{target_token: 128}), SamplingParams(temperature0.0), SamplingParams(temperature0.0, extra_args{target_token: 67}), SamplingParams(temperature0.0), ] llm LLM(modelfacebook/opt-125m, logits_processors[DummyLogitsProcessor]) outputs llm.generate(prompts, sampling_params_list)示例二包装请求级 processorcustom_req.py如果你的 processor 是按单请求粒度编写的——比如经典的f(output_ids, logits) - logits签名——直接塞给批级接口并不合适。examples/features/logits_processor/custom_req.py 演示了如何用AdapterLogitsProcessor基类AdapterLogitsProcessor 实现把它适配为批级 processorpython 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(ftarget_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_idsapply时逐行取logits[req_idx]交给对应请求的 processor若返回了新张量则回填到原行。由于 partial 持有输出 token 列表的引用processor 每步看到的output_ids始终是该请求截至当前的完整生成序列。示例三需要引擎配置的请求级包装custom_req_init.pyexamples/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( ftarget_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_tokensTPU 平台暂不支持当前 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_logitsprocs、AdapterLogitsProcessorvllm/v1/sample/logits_processor/builtin.py内建 processor 与process_dict_updates工具vllm/v1/sample/logits_processor/state.pyBatchUpdateBuilder与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),仅供参考
返回列表