vLLM Sampling Mask(Distribution Replay)深入解析:让 RL 训练中的 π_old 与 π_θ 共享同一个截断分布
【免费下载链接】vllmA high-throughput and memory-efficient inference and serving engine for LLMs项目地址: https://gitcode.com/GitHub_Trending/vl/vllm
Sampling Mask(采样掩码,官方文档也称 Distribution Replay / 分布重放)是 vLLM 为强化学习(RL)训练场景提供的一项引擎级特性。在 GRPO 等 rollout 采样中,top-k/top-p 截断会造成"采样器实际使用的截断分布"与"训练时全词表 softmax"之间的系统性不匹配,进而破坏重要性采样比 π_θ/π_old 的数学一致性;本特性通过返回每个生成步真正存活下来的 token 集合,让训练侧能在与 rollout 完全相同的支撑集(support)上归一化。读完本文,你将掌握该特性的启用参数、前置约束、底层实现链路,以及如何在 RL 训练框架中用 mask 正确计算当前策略的对数概率。
为什么要 Sampling Mask:截断分布与全词表 softmax 的错位
在基于 RL(如 GRPO)的 rollout 阶段,vLLM 采样器实际执行的是截断后的分布:logits 经过温度缩放、min-p、top-k/top-p 过滤后,被排除的 token 的 logits 会被置为-inf,采样只发生在幸存者集合中。然而,训练阶段计算对数概率时通常使用全词表 softmax。
二者叠加会产生一个微妙但致命的问题:π_old(旧策略)真正采样的动作空间是"截断后的核(nucleus)",而π_old(a|s)数值上却来自"全词表归一化";π_θ(当前策略)同理。两个策略在 importance sampling 中面对的动作子空间不一致,重要性比率失真,会破坏训练稳定性。
本特性对应 DeepSeek-V3.2 技术报告中描述的Keep Sampling Mask(保留采样掩码)策略:把 rollout 采样时由 top-k/top-p 截断产生的掩码保存下来,训练时把同一掩码套用到π_θ上,使新旧策略共享完全相同的动作子空间。报告指出,将 top-p 采样与 Keep Sampling Mask 结合,能有效在 RL 训练中保持语言一致性(语言流畅性不被破坏)。vLLM 的这一实现,正是把该策略从论文落到工程实践。
快速开始:三种接入方式
1. OpenAI 兼容服务 CLI
vllm serve <model> \ --return-sampling-mask \ --logprobs-mode processed_logprobs--return-sampling-mask是引擎级开关(默认关闭),--logprobs-mode processed_logprobs保证返回的 logprobs 是在截断核上归一化而不是全词表。两者在 参数解析入口 注册,最终落在 ModelConfig 的两个字段上:
return_sampling_mask: bool = False # """Whether to return the post-processing token support for each sample.""" logprobs_mode: LogprobsMode = "raw_logprobs" # """Indicates the content returned in the logprobs and prompt_logprobs."""2. 离线 Python API(LLM类)
from vllm import LLM, SamplingParams llm = LLM(model, return_sampling_mask=True, logprobs_mode="processed_logprobs") output = llm.generate( "The capital of France is", SamplingParams(temperature=1.0, top_k=50, top_p=0.95, logprobs=1), ) mask = output[0].outputs[0].sampling_mask # mask.token_ids: [[187, 326, 512], [42, 88], ...] # mask.token_ids[i] = token IDs in the sampling support for generated token iCompletionOutput新增的sampling_mask字段在 vllm/outputs.py 中是一个独立 dataclass:
@dataclass class SamplingMask: """Per-token sampling support sets aligned with completion token IDs. Each inner list contains the vocabulary token IDs that survived top-k / top-p / min-p filtering for the corresponding generated token. """ token_ids: list[list[int]]也就是说mask.token_ids是一个list[list[int]]:外层下标对应生成的每个 token,内层是该 token 生成时刻经过 top-k/top-p/min-p 过滤后仍存活的词表 token ID 集合。
3./inference/v1/generateHTTP 端点
掩码同样可以通过 scale-out 的 token-in-token-out 生成端点获取,协议定义 与 服务适配 会把SamplingMask.token_ids映射为响应字段:
{ "choices": [{ "token_ids": [187, 42, 303], "sampling_mask": [[187, 326, 512], [42, 88], [303, 11, 22]], "finish_reason": "stop" }] }注意这里sampling_mask的每一行严格对应token_ids中的每个已生成 token,逐位对齐,训练侧可以据此逐位置重建掩码。
前置要求与参数约束
启用前请先核对下表所列的四个必要条件:
| Requirement | Reason |
|---|---|
--return-sampling-mask | Engine-level opt-in(同时会禁用 FlashInfer 融合采样器) |
--logprobs-mode processed_logprobs | 返回的 logprobs 必须在截断核上归一化,而非全词表 |
temperature > 0 | Greedy(贪心)没有截断分布可言,掩码无意义 |
top_k > 0 | 约束掩码尺寸;纯 top-p 可能产生接近词表大小的掩码 |
| Model Runner V2 | 掩码需要异步 D2H(GPU→CPU)拷贝流水线支撑 |
引擎级配置校验
上述约束并非只写在文档里,引擎在VLLMConfig.__post_init__阶段通过_verify_sampling_replay_config做硬校验(vllm/config/vllm.py):一旦设置了return_sampling_mask,以下组合会在启动时直接抛ValueError:
- 非 Model Runner V2:报错 "sampling distribution replay requires Model Runner V2";
- Speculative decoding:不支持投机解码;
- Diffusion 模型:不支持扩散模型;
- 引擎级自定义 logits processors(
--logits-processors):不支持; logprobs_mode不是processed_logprobs:报错并要求显式设置,理由是"返回的 logprobs 必须与 sampling mask 在同一个截断核上归一化"。
此外,return_sampling_mask还与batch-sharded sampling不兼容:在 vllm/config/vllm.py 的 batch-sharded 采样可行性检查中,该组合会被列为 blocker(gather_sampler_output()不会转发SamplingMaskTensors,掩码会以None返回)。
请求级校验
除引擎级约束外,InputProcessor 在每个请求进入时还会校验SamplingParams:
temperature <= 0→ 报错 "sampling distribution replay requires temperature > 0"(greedy 没有截断分布);top_k <= 0→ 报错要求top_k > 0,理由注释写得很清楚:需要用它约束掩码尺寸、降低传输开销并避免潜在 OOM。
纯 top-p 之所以被排除,是因为最坏情况下幸存集合会膨胀到接近全词表,导致每步掩码的数据量与传输成本失控。
logprobs_mode:四种模式与"核上归一化"
--logprobs-mode支持四种取值(见 ModelConfig.logprobs_mode 注释),理解它们的区别是正确使用本特性的前提:
| 模式 | 返回内容 |
|---|---|
raw_logprobs | 未经过任何 logit processors(如 bad words、惩罚)处理前的 logprobs |
processed_logprobs | 应用全部 processors(含温度、top-k/top-p)之后的 logprobs |
raw_logits | 处理前的 logits |
processed_logits | 处理后的 logits |
当采样掩码开启时,采样器要求使用processed_logprobs:此时log_softmax是在processed logits(被过滤 token 为-inf)上计算的,softmax 的分母只覆盖截断核本身。这也意味着π_old(a|s)(rollout 侧、由 vLLM 产生)天然就是核上归一化的对数概率,与掩码代表的动作子空间一致。
一个实现细节(对理解不产生歧义,但值得注意):在 logits 分支,compute_topk_scores会根据logprobs_mode决定对原始 logits 还是处理后的 logits 取 top-k(采样器内实现);而 prompt token 不经过采样 processors,因此raw_*与processed_*对 prompt logprobs 而言结果相同(字段注释)。
工作原理:从 logits 到list[list[int]]的四步流水线
Step 1:应用全部 logit processors 并做 top-k/top-p 过滤
采样前,Sampler.sample 会按序在 logits 上原位施加:logit bias → 各类 penalty → bad words → thinking budget(若有)→ 温度 → min-p → top-k/top-p。top-k/top-p 过滤会把被排除 token 的 logits 置为-inf(processed_logits),幸存者保持有限值——这正是掩码的判据。
Step 2:用"有限 logits"找出幸存 token
采样完成后,若self.return_sampling_mask为真,Sampler.call会以本 batch 内校验过的最大 top-k 为宽度上限调用SamplingMaskTensors.from_logits。文档描述这一判据是torch.isfinite(processed_logits);实际落地的 Triton kernel_compact_sampling_mask_kernel(vllm/v1/worker/gpu/sample/output.py)则执行等价的keep = (logits > -inf) & (logits < inf)判据(NaN 因比较恒假也被排除)。
该 kernel 对每个请求行生成三种产物(SamplingMaskTensors):
token_ids:[num_requests, max_num_kept]的紧凑 ID 缓冲,仅存放每行前max_num_kept个存活 token;packed_mask:[num_requests, ceil(vocab_size / 8)]的逐位打包 bitmask;counts:每行存活 token 的精确数量。
其中紧凑 ID 宽度受双重约束:max_num_kept = min(top_k, vocab_size, MAX_COMPACT_SUPPORT),而MAX_COMPACT_SUPPORT = 2048([vllm/v1/worker/gpu/sample/output.py#L73-L74])。当某行存活 token 数超过该宽度时,以逐位打包的 bitmask 作为精确兜底编码——这保证了引擎不会因宽支撑而丢精度或内存失控。
Step 3:随采样 token 一起异步 D2H 拷贝
掩码是 GPU 张量,需要跨设备传输。这里走的是 Model Runner V2 的异步拷贝流(async D2H copy pipeline):在 async_utils.py 的 AsyncOutput 中,sampling_mask_tensors通过.to_cpu_nonblocking()在独立 CUDA copy stream 上与采样 token、logprobs 一同下发;源码特别注释了必须保留对 GPU 张量的引用,因为拷贝发生在与张量创建不同的流上。之后在get_output()中执行copy_event.synchronize()并把SamplingMaskTensors转成 CPU 侧的list([vllm/v1/worker/gpu/async_utils.py#L180-L183])。
Step 4:按请求合并、对齐并组包
每个解码步产出的掩码片段并不直接属于某个请求,需要经过调度与输出处理两级归并:
- 调度器持有
return_sampling_mask开关([vllm/v1/core/sched/scheduler.py#L381]),在把采样结果写回请求状态时用sampling_masks.slice_request(...)把掩码切分到对应请求([vllm/v1/core/sched/scheduler.py#L2111-L2131]); - 输出处理器为每个在途请求累积
sampling_mask_chunks([vllm/v1/engine/output_processor.py#L187]),当请求完成时把它们合并成一个SamplingMask([vllm/v1/engine/output_processor.py#L420-L436]),最终呈现为按生成位置对齐的list[list[int]]。
RL 训练侧使用:如何用 mask 计算 π_θ
训练侧计算重要性比率π_θ/π_old需要两个量,本特性恰好各解决一半。
π_old(a|s)——旧策略在截断核上归一化的对数概率:由 vLLM 在设置--logprobs-mode processed_logprobs后直接返回。因为log_softmax在 processed logits(被过滤 token 为-inf)上计算,分母只包含截断核,rollout 侧得到的每个 token 的 logprob 天然与采样真实分布一致,无需任何额外处理。
π_θ(a|s)——当前策略在同一个核上的对数概率:vLLM 不能替训练框架算 π_θ(π_θ 由训练侧正在更新的模型前向产生),因此框架需要用返回的 mask 自行实现"核上归一化":
# mask_ids: list[int], the sampling support for this token # logits: the training model's raw logits for this position keep = torch.zeros(vocab_size, dtype=torch.bool) keep[mask_ids] = True masked_logits = logits.masked_fill(~keep, float("-inf")) log_prob = log_softmax(masked_logits)[sampled_token_id]要点是:mask_ids必须使用 rollout 时由 vLLM 返回的掩码(即 π_old 的截断集合),并原样套用到 π_θ 上;两侧在同一 token 集合上做log_softmax,重要性比率π_θ/π_old的分母与支撑集严格一致,importance sampling 才站得住。
局限性与代价
从源码可以明确看到该特性的两项工程代价,使用前需要评估:
- 引擎级 flag 的全局代价:
--return-sampling-mask一旦开启,Sampler.init中use_flashinfer = not return_sampling_mask and flashinfer_sampler_supported()会全局禁用 FlashInfer 融合采样器。即使某个请求根本不需要掩码,也要走 PyTorch 采样路径并付出相应开销。对需要高吞吐的纯在线推理服务,建议与 RL 训练流量分开部署。 - 无流式支持:掩码只在最终响应中返回,不会出现在中间 streaming chunk 中。RL 训练通常以离线整段生成方式消费 rollout,一般不受影响;但任何依赖逐 token 流式拿到掩码的场景当前都无法满足。
此外,掩码按位置与生成 token 严格对齐,属于"逐位置"数据而非"逐请求标量",请求越长累积的传输与内存占用越大——这正是引擎强制top_k > 0、并用MAX_COMPACT_SUPPORT = 2048+ bitmask 兜底来控制数据量的原因。
小结
Sampling Mask 是一个精准服务 RL 训练管线的特性:它不改变 vLLM 的采样行为,而是把"采样器实际用过的截断支撑集"原样暴露给训练侧,让 π_old 与 π_θ 在完全相同的 token 集合上归一化。配合--logprobs-mode processed_logprobs、temperature > 0、top_k > 0与 Model Runner V2 四个前提,并在引擎配置校验(vllm/config/vllm.py)、请求级校验(vllm/v1/engine/input_processor.py)、Triton 掩码 kernel(vllm/v1/worker/gpu/sample/output.py)与异步 D2H 拷贝(vllm/v1/worker/gpu/async_utils.py)等源码环节都能找到一一对应的实现佐证。如果你正在搭建基于 vLLM 的 GRPO / RLHF 数据管线,需要为 rollout 与训练之间消除分布错位,这个特性就是现成的标准答案。
【免费下载链接】vllmA high-throughput and memory-efficient inference and serving engine for LLMs项目地址: https://gitcode.com/GitHub_Trending/vl/vllm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考