vLLM Sampling Mask(Distribution Replay)深入解析:让 RL 训练中的 π_old 与 π_θ 共享同一个截断分布
2026/9/7 8:26:17 网站建设 项目流程

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 i

CompletionOutput新增的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,逐位对齐,训练侧可以据此逐位置重建掩码。

前置要求与参数约束

启用前请先核对下表所列的四个必要条件:

RequirementReason
--return-sampling-maskEngine-level opt-in(同时会禁用 FlashInfer 融合采样器)
--logprobs-mode processed_logprobs返回的 logprobs 必须在截断核上归一化,而非全词表
temperature > 0Greedy(贪心)没有截断分布可言,掩码无意义
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 置为-infprocessed_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:按请求合并、对齐并组包

每个解码步产出的掩码片段并不直接属于某个请求,需要经过调度与输出处理两级归并:

  1. 调度器持有return_sampling_mask开关([vllm/v1/core/sched/scheduler.py#L381]),在把采样结果写回请求状态时用sampling_masks.slice_request(...)把掩码切分到对应请求([vllm/v1/core/sched/scheduler.py#L2111-L2131]);
  2. 输出处理器为每个在途请求累积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.inituse_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_logprobstemperature > 0top_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),仅供参考

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

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

立即咨询