vLLM V1 Sample 模块超深度架构分析 — Part 1: 架构总览与核心采样器

分析范围: vllm/v1/sample/ 目录全部源码(15个Python文件,约4,500行)
分析日期: 2026-05-25


目录


第一章 模块定位与全局架构

1.1 业务职责与功能定位

vllm/v1/sample 模块是 vLLM V1 架构中推理采样阶段的核心实现。其业务职责为:

  1. Logits后处理:对模型输出的原始logits施加一系列变换(logits处理器、惩罚项、温度缩放、top-k/top-p过滤)
  2. Token采样决策:从处理后的概率分布中采样下一个token(贪心采样或随机采样)
  3. 日志概率计算:计算采样token及top-N token的logprobs
  4. 投机解码采样:实现基于拒绝采样的推测解码(speculative decoding)验证与修正
  5. 思考预算控制:管理推理模型的thinking token预算,强制结束思考阶段

功能定位一句话总结:sample模块是从"模型输出logits"到"最终采样token"之间的完整决策管线。

1.2 在系统中的位置

Sample内部

Sampler

LogitsProcessors

TopKTopPSampler

Penalties

BadWords

RejectionSampler

Triton Kernels

ThinkingBudget

vLLM V1 推理引擎

logits

SamplerOutput

Scheduler
调度器

Worker
工作进程

GPUModelRunner
模型执行器

Model Forward
模型前向推理

🎯 Sample Module
采样模块

OutputProcessor
输出处理器

RequestHandler
请求处理

上游依赖:

  • GPUModelRunner.execute_model() → 调用 Sampler.forward(),传入模型输出的logits和SamplingMetadata
  • SamplingMetadata 由 GPUInputBatch 在每步构建

下游输出:

  • SamplerOutput(包含 sampled_token_ids + logprobs_tensors)→ 返回给 GPUModelRunner
  • RejectionSampler.forward() → 投机解码路径的替代采样入口

1.3 模块全景架构图

v1/sample/

logits_processor/

ops/

__init__.py
build_logitsprocs
AdapterLogitsProcessor

sampler.py
Sampler

metadata.py
SamplingMetadata

topk_topp_sampler.py
TopKTopPSampler

penalties.py
apply_all_penalties

bad_words.py
apply_bad_words

topk_topp_triton.py
Triton Kernel

rejection_sampler.py
RejectionSampler

interface.py
LogitsProcessor
BatchUpdate
MoveDirectionality

state.py
BatchUpdateBuilder
LogitsProcessors

builtin.py
MinP / LogitBias / MinTokens

thinking_budget_state.py
ThinkingBudgetStateHolder

logprobs.py
batched_count_greater_than

1.4 文件结构与代码量统计

文件路径行数核心类/函数职责
sampler.py425Sampler主采样器入口,协调所有采样步骤
metadata.py55SamplingMetadata采样参数数据容器
ops/topk_topp_sampler.py458TopKTopPSampler + 辅助函数top-k/top-p过滤+随机采样
ops/topk_topp_triton.py1,058_topk_topp_kernel + apply_top_k_top_p_tritonTriton加速的top-k/top-p
ops/penalties.py57apply_all_penalties惩罚项应用
ops/bad_words.py57apply_bad_words / apply_bad_words_with_drafts禁词屏蔽
ops/logprobs.py27batched_count_greater_thanlogprobs辅助计算
rejection_sampler.py921RejectionSampler + Triton kernels投机解码拒绝采样
thinking_budget_state.py528ThinkingBudgetStateHolder思考token预算管理
logits_processor/__init__.py357build_logitsprocs / AdapterLogitsProcessorlogits处理器构建与适配器
logits_processor/interface.py106LogitsProcessor / BatchUpdate处理器抽象接口
logits_processor/state.py165BatchUpdateBuilder / LogitsProcessors批次更新构建器
logits_processor/builtin.py332MinP / LogitBias / MinTokens内置logits处理器
合计~4,568

第二章 核心数据结构深度解析

2.1 SamplingMetadata — 采样元数据

@dataclass
class SamplingMetadata:
    temperature: torch.Tensor | None          # [batch_size] 每请求温度参数
    all_greedy: bool                           # 全批次是否均为贪心采样
    all_random: bool                           # 全批次是否均为随机采样
    
    top_p: torch.Tensor | None                 # [batch_size] top-p阈值
    top_k: torch.Tensor | None                 # [batch_size] top-k值
    
    generators: dict[int, torch.Generator]     # req_index → 随机数生成器(可复现性)
    
    max_num_logprobs: int | None               # None=不计算, 0=仅采样token, N=top-N
    
    no_penalties: bool                         # 全批次是否无需惩罚
    
    prompt_token_ids: torch.Tensor | None      # [batch_size, max_prompt_len] prompt token ids
    frequency_penalties: torch.Tensor           # [batch_size] 频率惩罚系数
    presence_penalties: torch.Tensor            # [batch_size] 存在惩罚系数
    repetition_penalties: torch.Tensor          # [batch_size] 重复惩罚系数
    
    output_token_ids: list[list[int]]           # 每请求已输出token id列表
    
    allowed_token_ids_mask: torch.Tensor | None # [max_batch, vocab_size] bool掩码
    
    bad_words_token_ids: dict[int, list[list[int]]]  # req_index → 禁词token序列
    
    logitsprocs: LogitsProcessors               # 已加载的logits处理器集合
    
    logprob_token_ids: dict[int, list[int]] | None   # req_index → 需计算logprob的特定token
    spec_token_ids: list[list[int]] | None           # 投机解码的draft token ids
    thinking_budget_state_holder: ThinkingBudgetStateHolder | None  # 思考预算状态

设计目的分析:

  • all_greedy / all_random:快速路径标记。当全批次均为同一模式时,可以跳过分支判断直接走优化路径
  • temperature: Tensor | None:使用Tensor而非list,因为需要在GPU上做批量温度缩放;None表示全贪心
  • generators: dict[int, Generator]:稀疏存储——只有需要可复现性的请求才提供generator,其余用默认随机种子
  • no_penalties: bool:提前判断可跳过惩罚计算的快速路径
  • output_token_ids: list[list[int]]:使用Python list而非Tensor,因为每个请求的输出长度不同(锯齿状),且需要频繁追加
  • allowed_token_ids_mask:2D bool掩码而非token id列表,因为需要在GPU上做批量masked_fill_
  • bad_words_token_ids: dict[int, ...]:稀疏存储,大多数请求没有禁词约束

SamplingMetadata

+temperature: Tensor|None

+all_greedy: bool

+all_random: bool

+top_p: Tensor|None

+top_k: Tensor|None

+generators: dict<int,Generator>

+max_num_logprobs: int|None

+no_penalties: bool

+prompt_token_ids: Tensor|None

+frequency_penalties: Tensor

+presence_penalties: Tensor

+repetition_penalties: Tensor

+output_token_ids: list<list<int>>

+allowed_token_ids_mask: Tensor|None

+bad_words_token_ids: dict

+logitsprocs: LogitsProcessors

+logprob_token_ids: dict|None

+spec_token_ids: list|None

+thinking_budget_state_holder: ThinkingBudgetStateHolder|None

LogitsProcessors

+argmax_invariant: list<LogitsProcessor>

+non_argmax_invariant: list<LogitsProcessor>

+all: Iterator

ThinkingBudgetStateHolder

+think_start_token_ids: list<int>

+think_end_token_ids: list<int>

+_state: dict<int,dict>

+mask: Tensor

+force_token_ids: Tensor

+has_tracked_requests() : bool

+sync_batch(batch_update)

+update_state(output_token_ids, spec_token_ids)

+apply_to_logits(logits, predict_bonus_token)

2.2 LogprobsMode 枚举与日志概率策略

LogprobsMode 是一个字符串字面量类型(来自 vllm.config.model),取值范围:

模式含义日志概率来源影响采样管线
"raw_logprobs"原始logprobs在惩罚/温度前计算 log(softmax(logits))最精确,但需提前计算
"raw_logits"原始logits在惩罚/温度前直接使用logits值不做softmax,近似logprobs
"processed_logits"处理后logits使用经过所有处理后的logits包含惩罚/温度/过滤效果
"processed_logprobs"处理后logprobs处理后logits再做log_softmax最接近实际采样分布

设计动机:不同场景对logprobs精度和性能有不同需求。raw_logprobs最精确但计算量最大;processed_*能反映实际采样分布但需要额外的kernel支持。

2.3 数据流向全景图

Step 7: 采样决策

Yes

No

Yes

No

all_greedy?

argmax(logits)

温度缩放 logits/temp

argmax不变处理器
MinP

Top-K / Top-P过滤

随机采样
Gumbel-Max / FlashInfer

temperature < ε?

随机采样结果

Step 1: Logprobs预处理

raw_logprobs模式:
log(softmax(logits))

raw_logits模式:
logits.clone().float32()

模型输出 logits
[batch_size, vocab_size]

Step 2: Float32转换
logits.to(float32)

Step 3: Allowed token IDs白名单
masked_fill_(mask, -inf)

Step 4: Bad words排除
logits[token_id] = -inf

Step 5: 非argmax不变处理器
MinTokens / LogitBias

Step 6: 惩罚项
Repetition / Frequency / Presence

Step 8: Gather logprobs
top-N + sampled token

SamplerOutput
(sampled_token_ids + logprobs_tensors)


第三章 Sampler 核心采样器逐行解析

3.1 类结构与初始化

class Sampler(nn.Module):
    """
    A layer that samples the next tokens from the model's outputs
    with the following steps in order:
    1. Compute logprobs (if requested)
    2. Convert logits to float32
    3. Apply allowed token ids whitelist
    4. Apply bad words exclusion
    5. Apply non-argmax-invariant logit processors (MinTokens, LogitBias)
    6. Apply penalties (Repetition, Frequency, Presence)
    7. Sample the next tokens (greedy or random with top-k/top-p)
    8. Gather top-N logprobs
    9. Return SamplerOutput
    """

Sampler

+topk_topp_sampler: TopKTopPSampler

+pin_memory: bool

+logprobs_mode: LogprobsMode

+forward(logits, sampling_metadata, predict_bonus_token, logprobs_mode_override) : SamplerOutput

+gather_specific_token_logprobs(logprobs, token_ids, batch_size) : LogprobsTensors

+apply_temperature(logits, temperature) : Tensor

+greedy_sample(logits) : Tensor

+sample(logits, sampling_metadata) : tuple

+compute_logprobs(logits) : Tensor

+gather_logprobs(logprobs, num, token_ids) : LogprobsTensors

+_combine_outputs_with_spec_tokens(output, spec) : list

+apply_logits_processors(logits, metadata, predict_bonus_token) : Tensor

+apply_penalties(logits, metadata) : Tensor

TopKTopPSampler

+logprobs_mode: LogprobsMode

+forward_native(logits, generators, k, p) : tuple

+forward_cuda(logits, generators, k, p) : tuple

+forward_cpu(logits, generators, k, p) : tuple

+forward_hip(logits, generators, k, p) : tuple

+forward_xpu(logits, generators, k, p) : tuple

__init__ 逐行解析:

def __init__(self, logprobs_mode: LogprobsMode = "raw_logprobs"):
    super().__init__()  # 继承nn.Module,使其可作为模型层使用
    # 实例化top-k/top-p采样器,传递logprobs_mode以决定是否可使用FlashInfer优化
    self.topk_topp_sampler = TopKTopPSampler(logprobs_mode)
    # 检测系统是否支持pin_memory(CUDA异步传输优化)
    self.pin_memory = is_pin_memory_available()
    # 存储logprobs模式,影响forward()中logprobs的计算时机
    self.logprobs_mode = logprobs_mode

设计要点:

  • Sampler 继承 nn.Module 而非普通类,是因为它需要在模型图中有明确位置,且 TopKTopPSampler 内部也继承 nn.Module
  • logprobs_mode 在构造时确定,但 forward() 支持 logprobs_mode_override 参数(投机解码路径需要不同模式)

3.2 forward() 主流程深度解析

def forward(
    self,
    logits: torch.Tensor,                     # [batch_size, vocab_size] 模型原始输出
    sampling_metadata: SamplingMetadata,       # 采样参数元数据
    predict_bonus_token: bool = False,         # 投机解码时是否预测bonus token
    logprobs_mode_override: LogprobsMode | None = None,  # logprobs模式覆盖
) -> SamplerOutput:

逐行解析:

# 确定本步使用的logprobs模式:优先使用覆盖值,否则使用默认值
logprobs_mode = logprobs_mode_override or self.logprobs_mode

# ===== Step 1: 计算raw logprobs(在任何变换之前)=====
# NOTE(woosuk): 使用原始logits(惩罚/温度之前)计算top-k logprobs
# 这与V0 sampler不同,V0使用变换后的logits
num_logprobs = sampling_metadata.max_num_logprobs
if num_logprobs is not None:  # 如果用户请求了logprobs
    if logprobs_mode == "raw_logprobs":
        # 计算log(softmax(logits))——最精确的原始日志概率
        raw_logprobs = self.compute_logprobs(logits)
    elif logprobs_mode == "raw_logits":
        # 直接使用logits值作为近似logprobs
        if logits.dtype == torch.float32:
            raw_logprobs = logits.clone()  # float32直接克隆,无需类型转换
        else:
            raw_logprobs = logits.to(torch.float32)  # 非float32需转换

关键设计决策:raw_logprobs 在惩罚/温度之前计算。这样做的原因是:

  1. 用户期望看到的是"模型对token的原始评价",而非经过人为调整后的值
  2. 与OpenAI API的logprobs语义对齐
  3. V0在变换后计算,导致logprobs包含惩罚效果,不符合预期
# ===== Step 2: 转换为float32 =====
logits = logits.to(torch.float32)
# 所有后续操作都在float32精度下进行,避免半精度的数值不稳定

# ===== Step 3-6: 应用logits处理器(包含allowed_token_ids、bad_words、惩罚等)=====
logits = self.apply_logits_processors(
    logits, sampling_metadata, predict_bonus_token
)

# ===== Step 7: 采样 =====
sampled, processed_logprobs = self.sample(logits, sampling_metadata)
# sampled: [batch_size] 采样出的token ids
# processed_logprobs: 如果logprobs_mode为processed_*,返回处理后的logprobs

# 如果sample()返回了processed_logprobs,覆盖之前计算的raw_logprobs
if processed_logprobs is not None:
    raw_logprobs = processed_logprobs

processed_logprobs覆盖机制:当logprobs_mode为processed_logits或processed_logprobs时,TopKTopPSampler会在top-k/top-p过滤后返回变换后的logits/logprobs,这些值比raw值更能反映实际采样分布。此时用processed值替换raw值。

# 转换采样结果为int64(long),确保与下游系统兼容
sampled = sampled.to(torch.int64)

# ===== Step 8: Gather logprobs =====
logprobs_tensors = None
if num_logprobs is not None:
    # 处理特定token logprobs请求(更高效,不需要遍历全词表)
    if sampling_metadata.logprob_token_ids is not None:
        logprobs_tensors = self.gather_specific_token_logprobs(
            raw_logprobs, sampling_metadata.logprob_token_ids, logits.shape[0]
        )
    else:
        # 标准路径:gather top-N logprobs + 采样token的logprob
        logprobs_tensors = self.gather_logprobs(
            raw_logprobs, num_logprobs, sampled
        )

# ===== Step 9: 构造输出 =====
return SamplerOutput(
    sampled_token_ids=sampled,     # [batch_size] int64
    logprobs_tensors=logprobs_tensors,  # LogprobsTensors | None
)

Yes

No

Yes

No

Yes

Yes

No

No

forward(logits, metadata)

logprobs_mode = override || default

max_num_logprobs≠None?

raw_logprobs → compute_logprobs(logits)
raw_logits → logits.clone()

logits → float32

apply_logits_processors(logits, metadata)

sample(logits, metadata)

sampled + processed_logprobs

processed_logprobs≠None?

raw_logprobs = processed_logprobs

sampled → int64

max_num_logprobs≠None?

logprob_token_ids≠None?

gather_specific_token_logprobs()

gather_logprobs(raw_logprobs, N, sampled)

logprobs_tensors = None

SamplerOutput

3.3 apply_logits_processors() 逻辑处理器调度

def apply_logits_processors(
    self,
    logits: torch.Tensor,               # [batch_size, vocab_size]
    sampling_metadata: SamplingMetadata,
    predict_bonus_token: bool = False,
) -> torch.Tensor:

逐行解析:

    # ===== Step 3: Allowed token IDs白名单 =====
    # 如果存在token白名单,将不在白名单中的token logits设为-inf
    if sampling_metadata.allowed_token_ids_mask is not None:
        # allowed_token_ids_mask: [max_batch, vocab_size], True表示"不允许"
        # 注意:mask的语义是"禁止",True=该token被屏蔽
        logits.masked_fill_(sampling_metadata.allowed_token_ids_mask, float("-inf"))
        # masked_fill_ 是原地操作,不分配新内存

设计注意:allowed_token_ids_mask 中 True 表示"该token不被允许"(反向语义)。这个命名容易混淆,但符合 masked_fill_ 的标准用法——True的位置被填充为指定值。

    # ===== Step 4: Bad words排除 =====
    if sampling_metadata.bad_words_token_ids:
        # bad_words_token_ids: {req_index: [[token_id1, token_id2, ...], ...]}
        # 每个禁词是一个token序列(可能多token),只屏蔽"即将生成的最后一个token"
        apply_bad_words(
            logits,
            sampling_metadata.bad_words_token_ids,
            sampling_metadata.output_token_ids,
        )
    # ===== Step 5: 非argmax不变logits处理器 =====
    # 这些处理器可能影响贪心采样的结果(改变argmax),因此必须在贪心采样前应用
    for processor in sampling_metadata.logitsprocs.non_argmax_invariant:
        logits = processor.apply(logits)
    # 典型的non-argmax-invariant处理器:
    # - MinTokensLogitsProcessor: 屏蔽EOS token,迫使模型生成更多token
    # - LogitBiasLogitsProcessor: 给特定token加偏置,可能改变argmax

    # ===== Thinking budget =====
    # 如果有思考预算追踪的请求,应用强制结束思考的逻辑
    holder = sampling_metadata.thinking_budget_state_holder
    if holder is not None and holder.has_tracked_requests():
        logits = holder.apply_to_logits(
            logits,
            predict_bonus_token=predict_bonus_token,
            spec_token_ids=sampling_metadata.spec_token_ids,
        )

    # ===== Step 6: 惩罚项 =====
    logits = self.apply_penalties(logits, sampling_metadata)

Yes

No

Yes

No

Yes

No

logits [B,V]

allowed_token_ids_mask?

masked_fill_(mask, -inf)

bad_words?

apply_bad_words()

non_argmax_invariant processors

MinTokens.apply()

LogitBias.apply()

thinking_budget?

holder.apply_to_logits()

apply_penalties()

processed logits

3.4 apply_penalties() 惩罚项应用

@staticmethod
def apply_penalties(
    logits: torch.Tensor,
    sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
    if sampling_metadata.no_penalties:
        return logits  # 快速路径:全批次无惩罚,直接跳过

    assert sampling_metadata.prompt_token_ids is not None
    # 有惩罚时必须有prompt_token_ids(重复惩罚需要知道哪些token在prompt中出现过)

    return apply_all_penalties(
        logits,
        sampling_metadata.prompt_token_ids,
        sampling_metadata.frequency_penalties,
        sampling_metadata.presence_penalties,
        sampling_metadata.repetition_penalties,
        sampling_metadata.output_token_ids,
    )

三种惩罚的区别:

惩罚类型公式效果argmax影响
Repetitionlogit > 0 → logit/rep_p; logit < 0 → logit*rep_p惩罚所有已出现token是(改变logit符号和大小)
Frequencylogit -= freq_p * count(token)按出现次数惩罚是(减去正值)
Presencelogit -= pres_p * (count > 0 ? 1 : 0)出现过就惩罚是(减去正值)

3.5 sample() 采样决策分支

def sample(
    self,
    logits: torch.Tensor,                # [batch_size, vocab_size] 处理后的logits
    sampling_metadata: SamplingMetadata,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    # 返回: (sampled_token_ids, processed_logprobs_or_None)

逐行解析:

    # ===== 7a: 贪心采样快速路径 =====
    if not sampling_metadata.all_random:
        # 如果有贪心请求(不是全随机),先做贪心采样
        argmax = logits.argmax(dim=-1)  # [batch_size] 每行最大值索引
        if sampling_metadata.all_greedy:
            # 全贪心:直接返回argmax结果
            return argmax, None

    # ===== 7b: 温度缩放 =====
    if sampling_metadata.temperature is not None:
        logits = self.apply_temperature(logits, sampling_metadata.temperature)
        # apply_temperature: logits /= temperature.unsqueeze(1)
        # temperature=0时在上方all_greedy路径已处理
    # ===== 7c: argmax不变logits处理器 =====
    # 这些处理器不改变argmax(贪心结果不变),只影响随机采样的概率分布
    for processor in sampling_metadata.logitsprocs.argmax_invariant:
        logits = processor.apply(logits)
    # 典型的argmax-invariant处理器:
    # - MinPLogitsProcessor: 过滤掉概率低于max_prob*min_p的token
    #   不影响argmax因为最大概率的token永远不会被过滤
    # ===== 7d-7e: Top-K/Top-P过滤 + 随机采样 =====
    sampled, maybe_processed_logprobs = self.topk_topp_sampler(
        logits=logits,
        generators=sampling_metadata.generators,
        k=sampling_metadata.top_k,
        p=sampling_metadata.top_p,
    )
    # 返回: (sampled_token_ids, processed_logprobs)
    # processed_logprobs仅当logprobs_mode为processed_*时有值
    # ===== 7f: 温度阈值判断 =====
    if not sampling_metadata.all_random:
        # 混合模式:需要根据温度决定每个请求用贪心还是随机
        # temperature < ε (1e-5) 的请求应使用贪心结果
        neq = sampling_metadata.temperature.ge(_SAMPLING_EPS)
        # neq[i] = True → 温度>=ε → 使用随机采样结果
        # neq[i] = False → 温度<ε → 使用贪心采样结果
        sampled = torch.where(neq, sampled, argmax)
        # 对每个请求,根据温度条件选择贪心或随机结果

    return sampled, maybe_processed_logprobs

No

Yes

Yes

No

Yes

No

sample(logits, metadata)

all_random?

argmax = logits.argmax(dim=-1)

apply_temperature(logits, temp)

all_greedy?

return (argmax, None)

argmax_invariant processors
(MinP)

topk_topp_sampler(logits, generators, k, p)

all_random?

return (sampled, processed_logprobs)

torch.where(temp≥ε, sampled, argmax)

return (merged, processed_logprobs)

3.6 compute_logprobs() 与 gather_logprobs()

compute_logprobs() — 静态方法:

@staticmethod
def compute_logprobs(logits: torch.Tensor) -> torch.Tensor:
    """计算 log(softmax(logits)) — 标准日志概率"""
    return logits.log_softmax(dim=-1, dtype=torch.float32)
    # 使用log_softmax而非log(softmax())避免数值溢出
    # 指定dtype=float32确保计算精度

gather_logprobs() — 完整解析:

def gather_logprobs(
    self,
    logprobs: torch.Tensor,      # [batch_size, vocab_size]
    num_logprobs: int,           # 需要gather的top-N数量
    token_ids: torch.Tensor,     # [batch_size] 采样出的token ids
) -> LogprobsTensors:
    # 1. 找出每行top-N个最大logprob及其token ids
    # topk返回: (values [B, N], indices [B, N])
    topk_logprobs, topk_token_ids = logprobs.topk(num_logprobs, dim=-1)
    
    # 2. 获取采样token自身的logprob
    # gather: 从每行选出token_ids指定位置的值
    token_logprobs = logprobs.gather(1, token_ids.unsqueeze(1)).squeeze(1)
    
    # 3. 构造LogprobsTensors
    return LogprobsTensors(
        logprobs=token_logprobs,         # [batch_size] 采样token的logprob
        token_ids=topk_token_ids,         # [batch_size, N] top-N token ids
        logprob_values=topk_logprobs,     # [batch_size, N] top-N logprob值
    )

3.7 gather_specific_token_logprobs() 特定token日志概率

def gather_specific_token_logprobs(
    self,
    logprobs: torch.Tensor,                      # [batch_size, vocab_size]
    logprob_token_ids: dict[int, list[int]],      # req_index → token_ids
    batch_size: int,
) -> LogprobsTensors:
    """为特定token计算logprobs——比全词表top-K更高效"""
    
    # 构建稀疏索引张量
    # 对于每个请求,只计算用户关心的那些token的logprob
    all_token_ids: list[list[int]] = []
    for i in range(batch_size):
        # 如果该请求指定了特定token ids,使用它们
        ids = logprob_token_ids.get(i, [])
        all_token_ids.append(ids)
    
    # 将锯齿状列表转为padded tensor
    # 使用0填充(vocab_size-1作为安全填充可能更好)
    max_len = max(len(ids) for ids in all_token_ids) if all_token_ids else 0
    if max_len == 0:
        return None  # 没有需要计算的token
    
    # 批量gather
    # indices: [batch_size, max_len] padded token ids
    indices = torch.zeros(batch_size, max_len, dtype=torch.long, device=logprobs.device)
    for i, ids in enumerate(all_token_ids):
        indices[i, :len(ids)] = torch.tensor(ids, dtype=torch.long)
    
    # gather操作: 从logprobs中按indices取出对应值
    gathered = logprobs.gather(1, indices)  # [batch_size, max_len]
    
    return LogprobsTensors(
        logprobs=None,
        token_ids=indices,
        logprob_values=gathered,
    )

第四章 TopKTopPSampler 采样算子深度解析

4.1 多平台策略模式

TopKTopPSampler 采用策略模式(Strategy Pattern),在 __init__ 中根据平台和配置动态绑定 self.forward 方法:

Yes

No

CUDA

Yes

Yes

No

No

CPU

RISCV/POWERPC

x86/ARM

XPU

Yes

No

ROCm

Yes

No

Other

TopKTopPSampler.__init__()

logprobs_mode是
processed_*?

self.forward = forward_native
(FlashInfer不支持返回中间值)

current_platform?

VLLM_USE_FLASHINFER_SAMPLER?

GPU compute capability支持?

self.forward = forward_cuda
(FlashInfer)

warning → forward_native

info → forward_native

CPU架构?

forward_native
(torch.compile在PPC上有bug)

self.forward = forward_cpu

VLLM_XPU_USE_SAMPLER_KERNEL?

self.forward = forward_xpu

forward_native

aiter_ops可用?

self.forward = forward_hip

forward_native

forward_native

策略选择决策表:

平台加速库条件forward方法
CUDAFlashInferVLLM_USE_FLASHINFER_SAMPLER=1 + GPU capability ≥ 8.0forward_cuda
CUDAPyTorch默认或FlashInfer不可用forward_native
CPUtorch.compilex86/ARM架构forward_cpu
CPUPyTorchRISC-V/PowerPCforward_native
ROCmaiteraiter_ops已安装forward_hip
XPU自定义kernelVLLM_XPU_USE_SAMPLER_KERNEL=1forward_xpu
其他PyTorch—forward_native

4.2 forward_native() PyTorch原生实现

def forward_native(
    self,
    logits: torch.Tensor,                          # [batch_size, vocab_size]
    generators: dict[int, torch.Generator],        # 稀疏随机数生成器
    k: torch.Tensor | None,                        # [batch_size] top-k值
    p: torch.Tensor | None,                        # [batch_size] top-p值
) -> tuple[torch.Tensor, torch.Tensor | None]:
    """PyTorch-native实现,适用于所有平台作为fallback"""
    
    # 1. 应用top-k和top-p过滤(内部选择Triton或PyTorch实现)
    logits = apply_top_k_top_p(logits, k, p)
    # logits中被过滤的位置被设为-inf
    
    # 2. 如果需要返回处理后的logits/logprobs
    logits_to_return = None
    if self.logprobs_mode == "processed_logits":
        logits_to_return = logits  # 直接返回处理后的logits
    elif self.logprobs_mode == "processed_logprobs":
        logits_to_return = logits.log_softmax(dim=-1, dtype=torch.float32)
        # 返回log(softmax(logits))
    
    # 3. 计算概率分布
    probs = logits.softmax(dim=-1, dtype=torch.float32)
    
    # 4. 随机采样(Gumbel-Max技巧)
    return random_sample(probs, generators), logits_to_return

Gumbel-Max技巧(random_sample内部实现):

  • 不使用 torch.multinomial(会导致CPU-GPU同步)
  • 而是生成指数分布噪声 q,然后 probs / q 的argmax等价于按概率采样
  • 数学证明:若 q_i ~ Exp(1),则 argmax_i(p_i / q_i) ~ Categorical(p)

4.3 forward_cuda() FlashInfer加速路径

def forward_cuda(
    self,
    logits: torch.Tensor,
    generators: dict[int, torch.Generator],
    k: torch.Tensor | None,
    p: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    """CUDA优化路径,使用FlashInfer的rejection sampling"""
    
    # 回退条件1: 没有top-k和top-p过滤
    # 回退条件2: 有per-request generators(FlashInfer 0.2.3+不支持)
    if (k is None and p is None) or generators:
        if generators:
            logger.debug_once(
                "FlashInfer 0.2.3+ does not support per-request generators. "
                "Falling back to PyTorch-native implementation."
            )
        return self.forward_native(logits, generators, k, p)
    
    # FlashInfer不支持返回中间logits/logprobs
    assert self.logprobs_mode not in ("processed_logits", "processed_logprobs"), (
        "FlashInfer does not support returning logits/logprobs"
    )
    
    # 确保logits是连续的(flex_attn/triton_attn fp32可能产生非连续张量)
    return flashinfer_sample(logits.contiguous(), k, p, generators), None

FlashInfer vs PyTorch-native 的性能差异:

特性FlashInferPyTorch-native
采样算法Rejection sampling(避免排序)Gumbel-Max(需softmax)
CPU-GPU同步无torch.multinomial有
per-request seed不支持支持
中间logprobs不支持支持
大词表性能O(V) 无排序O(V log V) 需排序

4.4 forward_cpu() CPU路径

def forward_cpu(
    self,
    logits: torch.Tensor,
    generators: dict[int, torch.Generator],
    k: torch.Tensor | None,
    p: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    """CPU优化路径"""
    
    # CPU上使用allow_cpu_sync=True的PyTorch实现
    # 区别:top-k-only时不需要排序整个词表
    logits = apply_top_k_top_p_pytorch(logits, k, p, allow_cpu_sync=True)
    
    logits_to_return = None
    if self.logprobs_mode == "processed_logits":
        logits_to_return = logits
    elif self.logprobs_mode == "processed_logprobs":
        logits_to_return = logits.log_softmax(dim=-1, dtype=torch.float32)
    
    # 无per-request generator时使用编译优化的快速采样
    if len(generators) != logits.shape[0]:
        return compiled_random_sample(logits), logits_to_return
        # compiled_random_sample: torch.compile包装的Gumbel-Max
    
    # 有per-request generator时逐请求生成指数噪声
    probs = logits.softmax(dim=-1, dtype=torch.float32)
    q = torch.empty_like(probs)
    q.exponential_()  # 批量生成Exp(1)噪声
    for i, generator in generators.items():
        q[i].exponential_(generator=generator)  # 覆盖有seed的行
    
    return probs.div_(q).argmax(dim=-1).view(-1), logits_to_return

4.5 forward_hip() ROCm/aiter路径

def forward_hip(
    self,
    logits: torch.Tensor,
    generators: dict[int, torch.Generator],
    k: torch.Tensor | None,
    p: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    """ROCm/aiter优化路径"""
    # FIXME: aiter_sampler存在精度问题,目前强制禁用
    DISABLE_AITER_SAMPLER = True
    
    if (k is None and p is None) or generators:
        return self.forward_native(logits, generators, k, p)
    
    assert self.logprobs_mode not in ("processed_logits", "processed_logprobs")
    
    if DISABLE_AITER_SAMPLER:
        return self.forward_native(logits, generators, k, p)
    
    return self.aiter_sample(logits, k, p, generators), None

aiter_sample() 内部支持三种路径:

def aiter_sample(self, logits, k, p, generators):
    use_top_k = k is not None
    use_top_p = p is not None
    
    if use_top_p and use_top_k:
        # 联合top-k+top-p: softmax → top_k_top_p_sampling_from_probs
        probs = logits.softmax(dim=-1, dtype=torch.float32).contiguous()
        return self.aiter_ops.top_k_top_p_sampling_from_probs(
            probs, None, *_to_tensor_scalar_tuple(k), *_to_tensor_scalar_tuple(p),
            deterministic=True
        ).view(-1)
    elif use_top_p:
        # 仅top-p: softmax → top_p_sampling_from_probs
        probs = logits.softmax(dim=-1, dtype=torch.float32).contiguous()
        return self.aiter_ops.top_p_sampling_from_probs(
            probs, None, *_to_tensor_scalar_tuple(p), deterministic=True
        ).view(-1)
    elif use_top_k:
        # 仅top-k: softmax → top_k_renorm → multinomial
        probs = logits.softmax(dim=-1, dtype=torch.float32).contiguous()
        renorm_probs = self.aiter_ops.top_k_renorm_probs(
            probs, *_to_tensor_scalar_tuple(k)
        )
        return torch.multinomial(renorm_probs, num_samples=1).view(-1)

4.6 forward_xpu() Intel XPU路径

def forward_xpu(
    self,
    logits: torch.Tensor,
    generators: dict[int, torch.Generator],
    k: torch.Tensor | None,
    p: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    """Intel XPU自定义kernel路径"""
    
    if generators:
        # XPU kernel不支持per-request generators
        return self.forward_native(logits, generators, k, p)
    
    # 分配输出张量
    random_sampled = torch.empty(
        logits.shape[0], dtype=torch.int64, device=logits.device
    )
    logits_to_return = None
    if self.logprobs_mode in ("processed_logits", "processed_logprobs"):
        logits_to_return = torch.empty_like(logits)
    
    # 获取XPU默认generator的seed/offset
    generator = torch.xpu.default_generators[logits.device.index]
    state = generator.get_state()
    seed, offset = state.view(torch.int64)
    seeds = torch.tensor([seed, offset], dtype=torch.int64, device="cpu")
    
    # top_k需要int64,但输入是int32——进行类型转换
    if k is not None:
        k = k.to(torch.int64)
    
    # 调用XPU自定义算子
    torch.ops.vllm.xpu_topk_topp_sampler(
        random_sampled, logits_to_return, logits, k, p, self.logprobs_mode, seeds
    )
    
    return random_sampled, logits_to_return

4.7 apply_top_k_top_p() 调度函数

def apply_top_k_top_p(
    logits: torch.Tensor, k: torch.Tensor | None, p: torch.Tensor | None
) -> torch.Tensor:
    """统一入口:选择Triton或PyTorch实现"""
    
    # 快速路径:无需过滤
    if p is None and k is None:
        return logits
    
    # Triton路径:batch_size >= 8且Triton可用
    if HAS_TRITON and logits.shape[0] >= 8:
        return apply_top_k_top_p_triton(logits, k, p)
    
    # PyTorch路径:小batch或无Triton
    return apply_top_k_top_p_pytorch(logits, k, p)

阈值8的设计理由:Triton kernel有启动开销(kernel launch + buffer分配),当batch_size < 8时,PyTorch排序可能更快。

4.8 apply_top_k_top_p_pytorch() 排序实现

def apply_top_k_top_p_pytorch(
    logits: torch.Tensor,       # [batch_size, vocab_size]
    k: torch.Tensor | None,     # [batch_size]
    p: torch.Tensor | None,     # [batch_size]
    allow_cpu_sync: bool = False,
) -> torch.Tensor:
    """基于排序的top-k/top-p实现"""
    
    # Top-k-only优化路径(CPU上避免排序整个词表)
    if p is None:
        if k is None:
            return logits
        if allow_cpu_sync:
            return apply_top_k_only(logits, k)
    
    # ===== 排序阶段 =====
    # 升序排序:logits_sort[i,j]是第i个请求第j小的logit
    logits_sort, logits_idx = logits.sort(dim=-1, descending=False)
    
    # ===== Top-K过滤 =====
    if k is not None:
        # top_k_mask: 保留的阈值 = 每行第(V-k)大的值
        top_k_mask = logits_sort.size(1) - k.to(torch.long)  # [B] 偏移量
        top_k_mask = logits_sort.gather(1, top_k_mask.unsqueeze(dim=1))
        # 将低于阈值的logit设为-inf
        top_k_mask = logits_sort < top_k_mask  # True = 需要屏蔽
        logits_sort.masked_fill_(top_k_mask, -float("inf"))
    
    # ===== Top-P过滤 =====
    if p is not None:
        # 从排序后的logits计算累积概率
        probs_sort = logits_sort.softmax(dim=-1)  # [B, V] 排序后的概率
        probs_sum = torch.cumsum(probs_sort, dim=-1, out=probs_sort)
        # probs_sum[i,j] = P(第j小及更小的token)
        
        # 保留累积概率 > 1-p 的token
        # 即: probs_sum <= 1-p 的token被屏蔽
        top_p_mask = probs_sum <= 1 - p.unsqueeze(dim=1)
        top_p_mask[:, -1] = False  # 确保至少保留1个token
        logits_sort.masked_fill_(top_p_mask, -float("inf"))
    
    # ===== 恢复原始顺序 =====
    # scatter: 将排序后的值放回原始位置
    return logits.scatter_(dim=-1, index=logits_idx, src=logits_sort)

Yes

No

Yes

No

logits [B,V]

升序排序
logits.sort(descending=False)

k≠None?

top-k过滤:
保留第(V-k)大的值及以上的

p≠None?

top-p过滤:
cumsum(softmax) > 1-p

scatter恢复原始顺序
logits.scatter_(idx, sort)

filtered logits [B,V]

Top-P过滤的数学原理:

  • 排序后从最小概率开始累积
  • cumsum ≤ 1-p 的token占比不超过 1-p,即"最不可能的 (1-p)*100%"
  • 屏蔽这些token后,剩余token的概率和 ≥ p
  • 设置 top_p_mask[:, -1] = False 保证至少一个token存活

4.9 apply_top_k_only() 无排序优化

def apply_top_k_only(
    logits: torch.Tensor,   # [batch_size, vocab_size]
    k: torch.Tensor,        # [batch_size]
) -> torch.Tensor:
    """Top-k-only优化:不需要排序整个词表
    
    注意:涉及GPU→CPU同步,对async scheduling性能有影响
    """
    # 找出不需要top-k的行(k == vocab_size = 不过滤)
    no_top_k_mask = k == logits.shape[1]
    k = k.masked_fill(no_top_k_mask, 1)  # 设为1以使gather操作有效
    
    # 获取每行的top-max_k值
    max_top_k = k.max()  # ⚠️ GPU→CPU同步点!
    top_k_mask = logits.topk(max_top_k, dim=1).values  # [B, max_top_k]
    
    # 找出每行的第k大的值(阈值)
    k_index = k.sub_(1).unsqueeze(1)  # 转为0-based索引
    top_k_mask = top_k_mask.gather(1, k_index.long())  # [B, 1] 阈值
    
    # 不需要过滤的行设为-inf阈值(实际不过滤任何token)
    top_k_mask.masked_fill_(no_top_k_mask.unsqueeze(1), -float("inf"))
    
    # 屏蔽低于阈值的logit
    return logits.masked_fill_(logits < top_k_mask, -float("inf"))

复杂度分析:

  • apply_top_k_top_p_pytorch:O(B * V * log V) — 需要排序
  • apply_top_k_only:O(B * V * max_k) — 只需topk,不排序
  • 当 max_k << V 时(典型情况 k=5~50, V=32000~128000),top-k-only显著更快

4.10 random_sample() Gumbel-Max采样

def random_sample(
    probs: torch.Tensor,                         # [batch_size, vocab_size]
    generators: dict[int, torch.Generator],      # 稀疏随机数生成器
) -> torch.Tensor:
    """Gumbel-Max采样:统计等价于Categorical采样,但避免CPU同步"""
    
    q = torch.empty_like(probs)
    
    # 批量生成Exp(1)噪声(无per-request seed的常见情况优化)
    if len(generators) != probs.shape[0]:
        q.exponential_()  # 全行批量生成
    
    # 对有seed的请求,逐个覆盖
    if generators:
        for i, generator in generators.items():
            q[i].exponential_(generator=generator)
    
    # argmax(probs / q) 等价于 Categorical(probs) 采样
    return probs.div_(q).argmax(dim=-1).view(-1)

Gumbel-Max技巧的数学证明:

设 p_i 是token i的概率,q_i ~ Exp(1) 独立指数分布。

则 argmax_i(p_i / q_i) 的分布恰好是 Categorical(p)。

证明:

  • p_i / q_i 等价于 log(p_i) - log(q_i)
  • log(q_i) = -G_i,其中 G_i ~ Gumbel(0,1)
  • 所以 argmax_i(log(p_i) + G_i) = Gumbel-Max采样

优势:

  1. 全GPU操作,无CPU同步
  2. 可批量处理
  3. 可复现(通过generator控制随机种子)

4.11 flashinfer_sample() FlashInfer采样

def flashinfer_sample(
    logits: torch.Tensor,
    k: torch.Tensor | None,
    p: torch.Tensor | None,
    generators: dict[int, torch.Generator],
) -> torch.Tensor:
    """FlashInfer采样:使用rejection sampling避免排序"""
    
    import flashinfer
    # 版本检查
    if version.parse(flashinfer.__version__) < version.parse("0.2.3"):
        raise ImportError("FlashInfer version >= 0.2.3 required")
    
    assert not (k is None and p is None)
    
    if k is None:
        # Top-p only
        probs = logits.softmax(dim=-1, dtype=torch.float32)
        next_token_ids = flashinfer.sampling.top_p_sampling_from_probs(
            probs, p, deterministic=True
        )
    elif p is None:
        # Top-k only
        probs = logits.softmax(dim=-1, dtype=torch.float32)
        next_token_ids = flashinfer.sampling.top_k_sampling_from_probs(
            probs, k, deterministic=True
        )
    else:
        # Top-k + Top-p
        next_token_ids = flashinfer.sampling.top_k_top_p_sampling_from_logits(
            logits, k, p, deterministic=True
        )
    
    return next_token_ids.view(-1)

FlashInfer rejection sampling 原理:

  1. 不对词表排序,而是生成均匀随机数u
  2. 对每个token,以概率p_i接受(p_i > u的阈值)
  3. 如果恰好一个token被接受则选中
  4. 否则重试
  5. 复杂度O(V)而非O(V log V)

4.12 compiled_random_sample() 编译优化

@torch.compile(dynamic=True)
def compiled_random_sample(logits: torch.Tensor) -> torch.Tensor:
    """torch.compile包装的Gumbel-Max采样,用于CPU路径
    
    dynamic=True: 允许动态形状,避免重复编译
    """
    probs = logits.softmax(dim=-1, dtype=torch.float32)
    q = torch.empty_like(probs)
    q.exponential_()
    return probs.div(q).argmax(dim=-1).view(-1)

为什么CPU需要 torch.compile:

  • CPU上没有FlashInfer
  • torch.multinomial 导致CPU同步
  • torch.compile 可以将 Gumbel-Max 融合为优化的CPU kernel
  • 注意:此路径不支持per-request generator(无seed参数)

第五章 Penalties惩罚算子深度解析

5.1 apply_all_penalties() 统一入口

def apply_all_penalties(
    logits: torch.Tensor,               # [batch_size, vocab_size]
    prompt_token_ids: torch.Tensor,     # [batch_size, max_prompt_len]
    presence_penalties: torch.Tensor,   # [batch_size]
    frequency_penalties: torch.Tensor,  # [batch_size]
    repetition_penalties: torch.Tensor, # [batch_size]
    output_token_ids: list[list[int]],  # 每请求已输出token列表
) -> torch.Tensor:
    """应用presence/frequency/repetition三种惩罚"""
    
    _, vocab_size = logits.shape
    
    # 将锯齿状的output_token_ids转为padded tensor
    output_tokens_t = _convert_to_tensors(output_token_ids, vocab_size, logits.device)
    
    # 异步调度中,不会应用惩罚的行可能包含-1占位符token id
    # 必须替换为有效token id,否则scatter操作会越界
    # NOTE(nick): 当前惩罚实现效率较低,后续会重做
    output_tokens_t.masked_fill_(output_tokens_t == -1, vocab_size)
    # vocab_size是合法的填充值,因为token id范围是[0, vocab_size)
    
    return apply_penalties(
        logits,
        prompt_token_ids,
        output_tokens_t,
        presence_penalties,
        frequency_penalties,
        repetition_penalties,
    )
    # 底层apply_penalties在vllm.model_executor.layers.utils中实现(CUDA kernel)

5.2 _convert_to_tensors() 张量转换

def _convert_to_tensors(
    output_token_ids: list[list[int]],  # 锯齿状已输出token
    vocab_size: int,
    device: torch.device,
) -> torch.Tensor:
    """将Python list的list转为padded tensor"""
    
    # 使用vocab_size作为padding值(因为token id ∈ [0, vocab_size))
    output_tokens_tensor = make_tensor_with_pad(
        output_token_ids,
        pad=vocab_size,           # padding值=词表大小
        device="cpu",             # 先在CPU上创建
        dtype=torch.int64,
        pin_memory=is_pin_memory_available(),  # pin_memory加速CPU→GPU传输
    )
    # 异步传输到GPU
    return output_tokens_tensor.to(device, non_blocking=True)

pin_memory优化:pin_memory=True 将CPU内存锁定(不被换出到磁盘),使 non_blocking=True 的GPU传输真正异步。

5.3 底层apply_penalties() CUDA内核逻辑

底层 apply_penalties() 位于 vllm/model_executor/layers/utils.py,实现为CUDA kernel,三种惩罚的计算逻辑:

Repetition Penalty(重复惩罚):

# 伪代码
for token_id in appeared_tokens:
    if logits[token_id] > 0:
        logits[token_id] /= repetition_penalty  # 正logit被缩小
    else:
        logits[token_id] *= repetition_penalty   # 负logit被放大
# 效果:降低已出现token的概率,repetition_penalty > 1时生效

Frequency Penalty(频率惩罚):

# 伪代码
for token_id, count in token_counts:
    logits[token_id] -= frequency_penalty * count
# 效果:按出现次数线性惩罚,出现越多惩罚越重

Presence Penalty(存在惩罚):

# 伪代码
for token_id in appeared_tokens:
    logits[token_id] -= presence_penalty  # 只惩罚一次
# 效果:只要出现过就惩罚固定量,不关心出现次数

logits + prompt_ids + output_ids

Repetition Penalty
正logit: ÷ rep_p
负logit: × rep_p

Frequency Penalty
logit -= freq_p × count

Presence Penalty
logit -= pres_p × (count > 0)

penalized logits


第六章 BadWords与Logprobs辅助算子

6.1 apply_bad_words() 禁词屏蔽

_SMALLEST_LOGIT = float("-inf")  # 屏蔽值,确保被屏蔽token概率为0

def _apply_bad_words_single_batch(
    logits: torch.Tensor,                  # [vocab_size] 单个请求的logits
    bad_words_token_ids: list[list[int]],  # [[t1,t2,...], ...] 禁词token序列
    past_tokens_ids: list[int],            # 已生成的token id序列
) -> None:
    """对单个请求应用禁词屏蔽(原地修改)"""
    
    for bad_word_ids in bad_words_token_ids:
        # 如果禁词序列比已生成序列+1还长,不可能匹配
        # "+1"是因为当前正在生成的token也参与匹配
        if len(bad_word_ids) > len(past_tokens_ids) + 1:
            continue
        
        prefix_length = len(bad_word_ids) - 1  # 前缀长度(不含最后一个token)
        last_token_id = bad_word_ids[-1]       # 需要屏蔽的token id
        
        # 获取已生成序列的最后prefix_length个token
        actual_prefix = past_tokens_ids[-prefix_length:] if prefix_length > 0 else []
        expected_prefix = bad_word_ids[:prefix_length]
        
        assert len(actual_prefix) == len(expected_prefix)
        
        # 前缀匹配:如果已生成token的尾部与禁词前缀匹配
        # 则屏蔽禁词的最后一个token
        if actual_prefix == expected_prefix:
            logits[last_token_id] = _SMALLEST_LOGIT

设计意图:禁词不是单个token,而是token序列(如"不"+“好”)。屏蔽逻辑是:

  1. 检查已生成序列的尾部是否匹配禁词的前缀
  2. 如果匹配,将禁词的最后一个token的logit设为-inf
  3. 这样模型就不会生成该token,从而避免生成完整的禁词序列
def apply_bad_words(
    logits: torch.Tensor,                              # [batch_size, vocab_size]
    bad_words_token_ids: dict[int, list[list[int]]],   # req_index → 禁词
    past_tokens_ids: list[list[int]],                  # 每请求已生成序列
) -> None:
    """批量应用禁词屏蔽"""
    for i, bad_words_ids in bad_words_token_ids.items():
        _apply_bad_words_single_batch(
            logits[i], bad_words_ids, past_tokens_ids[i]
        )

6.2 apply_bad_words_with_drafts() 投机采样禁词

def apply_bad_words_with_drafts(
    logits: torch.Tensor,                              # [num_tokens, vocab_size]
    bad_words_token_ids: dict[int, list[list[int]]],   # req_index → 禁词
    past_tokens_ids: list[list[int]],                  # 每请求已生成序列(含draft tokens)
    num_draft_tokens: list[int],                       # 每请求的draft token数量
) -> None:
    """投机解码场景的禁词屏蔽
    
    在投机解码中,logits形状为[num_tokens, vocab_size],
    其中num_tokens = sum(draft_tokens_per_request)
    每个请求可能有多行logits(对应多个draft位置)
    """
    start_idx = 0
    remaining = len(bad_words_token_ids)
    for i, n in enumerate(num_draft_tokens):
        if (bad_words_ids := bad_words_token_ids.get(i)) is not None:
            # 对该请求的每个draft位置都应用禁词屏蔽
            for draft_idx in range(start_idx, start_idx + n):
                _apply_bad_words_single_batch(
                    logits[draft_idx],
                    bad_words_ids,
                    past_tokens_ids[draft_idx],
                )
            remaining -= 1
            if not remaining:
                break  # 提前退出:所有有禁词的请求已处理完
        start_idx += n

投机解码禁词的特殊性:

  • 常规采样:每个请求1行logits → 检查1次
  • 投机解码:每个请求N行logits(N=draft tokens数)→ 检查N次
  • past_tokens_ids[draft_idx] 需要包含到该位置为止的所有已生成token

6.3 batched_count_greater_than() 日志概率辅助

@torch.compile(backend=current_platform.simple_compile_backend)
def batched_count_greater_than(
    x: torch.Tensor,      # [batch_size, n_elements]
    values: torch.Tensor,  # [batch_size, 1]
) -> torch.Tensor:
    """统计每行中大于等于给定值的元素数量
    
    用于logprobs计算中确定token排名
    """
    torch._check(x.shape[0] >= 1)
    torch._check(x.shape[0] == values.shape[0])
    return (x >= values).sum(-1)  # [batch_size]

为什么用 torch.compile:

  • 不编译时,(x >= values).sum(-1) 会创建中间bool张量,导致内存翻倍
  • torch.compile 将其融合为单个kernel,避免中间内存分配
  • 特别重要:当 x 是 [batch, 128000](大词表)时

附录A 采样管线完整时序图

ThinkingBudget BadWords Penalties TopKTopPSampler LogitsProcessors Sampler GPUModelRunner ThinkingBudget BadWords Penalties TopKTopPSampler LogitsProcessors Sampler GPUModelRunner Step 1: Compute raw logprobs Step 2: Float32 cast Step 3: Allowed token IDs Step 4: Bad words Step 5: Non-argmax-invariant processors MinTokens: mask EOS tokens LogitBias: add biases Thinking budget Step 6: Penalties Step 7: Sample MinP: filter low-prob tokens Step 8: Gather logprobs forward(logits, metadata) compute_logprobs(logits) [if requested] logits.to(float32) masked_fill_(allowed_mask, -inf) apply_bad_words(logits, bad_words, output_ids) logits modified in-place non_argmax_invariant.apply(logits) processed logits apply_to_logits(logits) forced end-thinking tokens apply_all_penalties(logits, penalties) penalized logits greedy_sample(logits) [if all_greedy] apply_temperature(logits, temp) argmax_invariant.apply(logits) forward(logits, generators, k, p) (sampled, processed_logprobs) torch.where(temp≥ε, sampled, argmax) gather_logprobs(raw_logprobs, N, sampled) SamplerOutput(sampled_token_ids, logprobs_tensors)

附录B 术语表

术语全称含义
Logits—模型输出的未归一化对数概率
LogprobsLog Probabilitieslog(softmax(logits)),归一化对数概率
Top-K—只保留概率最高的K个token
Top-P (Nucleus)—保留概率累积和≥P的最少token
Min-P—过滤概率低于 max_prob×min_p 的token
Gumbel-MaxGumbel-Max Trick用指数分布噪声实现分类采样的技巧
Repetition Penalty—对已出现token施加惩罚,防重复
Frequency Penalty—按出现次数线性惩罚已出现token
Presence Penalty—出现过就惩罚固定量
Argmax-Invariant—logit处理器不改变argmax(贪心采样结果不变)
Pin MemoryPinned Memory锁定CPU内存不被换出,加速CPU→GPU传输
Spec DecodeSpeculative Decoding投机解码,用小模型预测+大模型验证
Rejection Sampling—拒绝采样,投机解码的核心验证算法
FlashInfer—高性能GPU采样库
Triton—OpenAI的GPU编程语言,用于自定义kernel
aiter—AMD ROCm的加速推理库
XPU—Intel GPU加速平台
Bonus Token—投机解码中所有draft token被接受后额外采样的token
Draft Token—投机解码中小模型预测的token
Recovered Token—拒绝采样中拒绝后重新采样的token
Thinking Budget—推理模型的思考token预算限制
CUDA Graph—CUDA操作图,用于减少kernel launch开销

附录I Sampler.forward() 完整伪代码与行号映射

以下为 sampler.py 中 Sampler.forward() 的完整执行路径,逐行标注代码行号与执行语义。

I.1 Logprobs预处理阶段(L62-L82)

L62: logprobs_mode = logprobs_mode_override or self.logprobs_mode
     # 确定本步logprobs计算模式
     # logprobs_mode_override 来自投机解码的覆盖需求

L63: num_logprobs = sampling_metadata.max_num_logprobs
     # 从元数据获取用户请求的logprobs数量
     # None = 不计算, 0 = 仅采样token, N = top-N

L65: if num_logprobs is not None:
     # 用户请求了logprobs → 需要在变换前保存原始值

L67:     if logprobs_mode == "raw_logprobs":
         # 模式1: 计算log(softmax(logits))作为原始logprobs
L68:         raw_logprobs = self.compute_logprobs(logits)
         # 实现: logits.log_softmax(dim=-1, dtype=torch.float32)
         # 选择log_softmax而非log(softmax())避免数值溢出
         # 指定float32确保半精度输入的计算精度

L70:     elif logprobs_mode == "raw_logits":
         # 模式2: 直接使用logits值作为近似logprobs
L71:         if logits.dtype == torch.float32:
L72:             raw_logprobs = logits.clone()
             # float32直接clone,无需类型转换
             # clone()创建独立张量,后续对logits的修改不影响raw_logprobs
L74:             raw_logprobs = logits.to(torch.float32)
             # 非float32(如bfloat16/float16)需转换精度
             # 这一步创建了新张量,后续修改不影响raw_logprobs

I.2 Float32转换(L84-L85)

L84: logits = logits.to(torch.float32)
     # 所有后续操作在float32精度下进行
     # 原因:
     # 1. softmax在float16下可能溢出(指数部分超过65504)
     # 2. 温度除法在低精度下数值不稳定
     # 3. 惩罚计算需要精确的加减操作
     # 4. logprobs计算需要log_softmax的数值稳定性
     # 注意: 如果logits已经是float32,to()是no-op

I.3 Logits处理器调度(L87-L89)

L87: logits = self.apply_logits_processors(
L88:     logits, sampling_metadata, predict_bonus_token
L89: )
     # 依次执行:
     # 1. allowed_token_ids_mask → masked_fill_(mask, -inf)
     # 2. bad_words → apply_bad_words()
     # 3. non_argmax_invariant processors (MinTokens, LogitBias)
     # 4. thinking_budget → holder.apply_to_logits()
     # 5. penalties → apply_all_penalties()
     # predict_bonus_token: 投机解码标记,影响thinking_budget的行为

I.4 采样决策(L91-L96)

L91: sampled, processed_logprobs = self.sample(logits, sampling_metadata)
     # sample()内部流程:
     # 1. all_greedy → 直接argmax
     # 2. 温度缩放: logits /= temperature
     # 3. argmax_invariant processors (MinP)
     # 4. top-k/top-p过滤 + 随机采样
     # 5. 温度阈值判断: temp < ε → 用贪心结果
     # 返回: (sampled_token_ids [B], processed_logprobs | None)

L93: if processed_logprobs is not None:
     # 当logprobs_mode为processed_*时,TopKTopPSampler返回
     # 经过温度/top-k/top-p处理后的logits/logprobs
     # 这些值更能反映实际采样的概率分布

L94:     raw_logprobs = processed_logprobs
     # 用处理后的值覆盖原始logprobs
     # 因为用户想看到的是"采样时的实际概率",而非"模型原始输出"

I.5 类型转换与Logprobs聚合(L97-L112)

L97: sampled = sampled.to(torch.int64)
     # 转为int64(long)确保与下游系统兼容
     # 下游: Scheduler、RequestHandler等使用long类型token id

L99: logprobs_tensors = None
L100: if num_logprobs is not None:
     # 用户请求了logprobs → 需要gather

L102:     if sampling_metadata.logprob_token_ids is not None:
         # 优化路径: 只计算特定token的logprobs
         # 比全词表top-K更高效(用户只关心几个token)
L103:         logprobs_tensors = self.gather_specific_token_logprobs(
L104:             raw_logprobs, sampling_metadata.logprob_token_ids, logits.shape[0]
L105:         )

L108:         logprobs_tensors = self.gather_logprobs(
L109:             raw_logprobs, num_logprobs, sampled
L110:         )
         # gather_logprobs内部:
         # 1. topk(num_logprobs) → top-N logprob + token_ids
         # 2. gather(token_ids) → 采样token自身的logprob
         # 3. 构造LogprobsTensors

L113: return SamplerOutput(
L114:     sampled_token_ids=sampled,
L115:     logprobs_tensors=logprobs_tensors,
L116: )

附录J TopKTopPSampler 多平台路径对比矩阵

特性forward_nativeforward_cudaforward_cpuforward_hipforward_xpu
平台通用NVIDIA CUDACPUAMD ROCmIntel XPU
加速库PyTorch/TritonFlashInfertorch.compileaiter自定义kernel
Top-K/Top-P实现apply_top_k_top_pFlashInferapply_top_k_top_p_pytorchaiter_samplexpu_topk_topp_sampler
采样方式Gumbel-MaxFlashInfer rejectionGumbel-Max (compiled)aiter sampling自定义
Per-request seed✅❌✅❌❌
Processed logprobs✅❌✅❌✅
Batch阈值无无无无无
回退条件—k=None & p=None / generatorsgeneratorsk=None & p=None / generators / DISABLEgenerators
算子复杂度O(V log V) 或 O(V) TritonO(V) rejectionO(V log V) 或 O(V×k)O(V)O(V)
数值精度float32float32float32float32float32

J.1 各平台采样算法对比

XPU Custom Kernel

xpu_topk_topp_sampler(logits, k, p, seeds)

单kernel完成top-k/top-p/采样

aiter Sampling (hip)

softmax(logits) → probs

aiter top_k_renorm_probs

torch.multinomial(renorm_probs)

top-p: aiter top_p_sampling_from_probs

FlashInfer Rejection (cuda)

softmax(logits) → probs

生成均匀随机数u

对每个token: 接受概率p_i

如果恰好1个接受 → 选中

否则重试

复杂度O(V),无需排序

Gumbel-Max (native/cpu)

softmax(logits) → probs

q = exponential_(probs.shape)

sampled = argmax(probs / q)

统计等价于Categorical(probs)

J.2 apply_top_k_top_p 分支逻辑详解

apply_top_k_top_p(logits, k, p)
├── k=None and p=None → return logits (无过滤)
├── HAS_TRITON and batch >= 8 → apply_top_k_top_p_triton(logits, k, p)
│   ├── Top-K only: 三元搜索找k_pivot → mask
│   ├── Top-P only: 三元搜索找p_pivot → mask  
│   └── Top-K + Top-P: 先K后P → mask
└── else → apply_top_k_top_p_pytorch(logits, k, p)
    ├── p=None → apply_top_k_only(logits, k) [if allow_cpu_sync]
    │   └── topk(max_k) → gather阈值 → masked_fill_
    └── p≠None → sort → top-k mask → top-p mask → scatter恢复

apply_top_k_top_p_pytorch 详细步骤:

Step 1: 升序排序
  logits_sort, logits_idx = logits.sort(dim=-1, descending=False)
  # logits_sort[i,j] = 第i行第j小的值
  # logits_idx[i,j] = 原始位置索引

Step 2: Top-K过滤
  threshold_index = V - k  # 保留k个最大值
  threshold_value = logits_sort.gather(1, threshold_index)
  mask = logits_sort < threshold_value
  logits_sort.masked_fill_(mask, -inf)

Step 3: Top-P过滤
  probs_sort = softmax(logits_sort)  # 排序后的概率
  cumprobs = cumsum(probs_sort)      # 累积概率
  mask = cumprobs <= 1 - p           # 保留累积概率 > 1-p 的
  mask[:, -1] = False                # 至少保留1个token
  logits_sort.masked_fill_(mask, -inf)

Step 4: 恢复原始顺序
  logits.scatter_(dim=-1, index=logits_idx, src=logits_sort)
  # 将排序后的值放回原始位置

附录K 惩罚项数学推导与实现细节

K.1 Repetition Penalty 重复惩罚

公式:

logit'[t] = logit[t] / rep_p    if logit[t] > 0
logit'[t] = logit[t] * rep_p    if logit[t] < 0
logit'[t] = logit[t]            if logit[t] = 0

其中 rep_p ≥ 1,t 是已出现过的token。

推导:

  • 正logit表示模型倾向于选择该token → 除以rep_p缩小倾向
  • 负logit表示模型倾向于避免该token → 乘以rep_p增强避免
  • 这种不对称处理确保了惩罚效果的方向性

为什么不用统一的减法:

logit'[t] = logit[t] - penalty  # 不好的设计
  • 减法对所有logit一视同仁,但正负logit的语义不同
  • 正logit减去固定量可能变为负值,过度惩罚
  • 除法/乘法保持了logit的正负语义

K.2 Frequency Penalty 频率惩罚

公式:

logit'[t] = logit[t] - freq_p × count(t)

其中 count(t) 是token t 在输出序列中出现的次数。

效果分析:

  • 每多出现一次,logit减少 freq_p 的量
  • 线性递减,出现次数越多惩罚越重
  • freq_p > 0:抑制重复
  • freq_p < 0:鼓励重复(不常用)

K.3 Presence Penalty 存在惩罚

公式:

logit'[t] = logit[t] - pres_p × I(t ∈ output)

其中 I(t ∈ output) 是指示函数,token出现过为1,否则为0。

与Frequency的区别:

  • Frequency: 按次数惩罚,出现2次比1次惩罚更重
  • Presence: 只看是否出现过,出现100次和1次惩罚相同
  • Presence更适合"避免某个话题",Frequency更适合"避免重复用词"

K.4 三种惩罚的交互效果

场景: token "the" 出现了5次,原始logit = 4.0
参数: rep_p = 1.2, freq_p = 0.3, pres_p = 0.5

Step 1: Repetition
  logit = 4.0 / 1.2 = 3.33  (正logit,除以rep_p)

Step 2: Frequency
  logit = 3.33 - 0.3 × 5 = 3.33 - 1.5 = 1.83

Step 3: Presence
  logit = 1.83 - 0.5 × 1 = 1.33

最终: 4.0 → 1.33,降低了2.67,幅度66.75%

原始 logit = 4.0

Repetition
4.0/1.2 = 3.33

Frequency
3.33-1.5 = 1.83

Presence
1.83-0.5 = 1.33

最终 logit = 1.33
降低66.75%


附录L BadWords 禁词过滤深度分析

L.1 多token禁词序列匹配算法

给定:
  past_tokens = [10, 20, 30, 40, 50]  # 已生成5个token
  bad_word = [30, 40, 999]             # 3-token禁词序列

匹配过程:
  prefix_length = len(bad_word) - 1 = 2  # 前缀长度(不含最后一个)
  last_token_id = 999                     # 需要屏蔽的token

  actual_prefix = past_tokens[-2:] = [40, 50]  # 最后2个token
  expected_prefix = bad_word[:2] = [30, 40]     # 禁词前缀

  actual_prefix == expected_prefix? 
  [40, 50] == [30, 40]? → False → 不屏蔽

  # 如果past_tokens = [10, 20, 30, 40, ...]:
  actual_prefix = [30, 40]
  expected_prefix = [30, 40]
  [30, 40] == [30, 40]? → True → 屏蔽 token 999

L.2 单token禁词的简化路径

给定:
  bad_word = [999]  # 单token禁词
  
  prefix_length = 0
  last_token_id = 999
  actual_prefix = []
  expected_prefix = []
  [] == [] → True → 总是屏蔽 token 999
  # 单token禁词等价于将该token加入黑名单

L.3 超长禁词的提前退出

给定:
  past_tokens = [10, 20]  # 只生成了2个token
  bad_word = [1, 2, 3, 4, 5]  # 5-token禁词

  len(bad_word) > len(past_tokens) + 1
  5 > 2 + 1 = 3 → True → continue (跳过)
  # 已生成的token不足以匹配禁词前缀
  # "+1" 是因为当前正在生成的token也参与匹配

L.4 投机解码中的BadWords应用

在投机解码中,每个请求有N个draft位置,每个位置需要独立检查禁词:

请求0: 2个draft tokens, bad_words = [[5, 6, 100]]
  draft位置0: past=[1, 2, 5, 6], 检查[5,6,100] → actual=[5,6]==[5,6] → 屏蔽100
  draft位置1: past=[1, 2, 5, 6, draft0_token], 再次检查

请求1: 3个draft tokens, 无bad_words
  跳过

提前退出: remaining=1, 处理完请求0后remaining=0, break

Yes

No

logits [num_tokens, V]
bad_words_token_ids
num_draft_tokens=[2,3,1]

start_idx=0, remaining=num_bad_words_reqs

遍历num_draft_tokens

请求0: n=2, has_bad_words

draft_idx 0: check bad_words

draft_idx 1: check bad_words

remaining-=1, remaining=0?

break提前退出

请求1: n=3, no bad_words

start_idx += 3


附录M batched_count_greater_than() 编译优化深度分析

M.1 为什么需要torch.compile

# 未编译版本:
def naive_count(x, values):
    return (x >= values).sum(-1)
# 问题:
# 1. x >= values 创建 [B, V] bool中间张量
# 2. sum(-1) 需要读取整个bool张量
# 3. 当V=128000时,中间张量=128000 bytes per row
# 4. 如果B=256,中间张量=32MB
# 5. 这32MB在下一帧前无法释放 → 内存压力

# 编译版本:
@torch.compile(backend=current_platform.simple_compile_backend)
def compiled_count(x, values):
    torch._check(x.shape[0] >= 1)
    torch._check(x.shape[0] == values.shape[0])
    return (x >= values).sum(-1)
# torch.compile效果:
# 1. 融合 x >= values 和 sum(-1) 为单个kernel
# 2. 中间bool结果不写入内存,直接在寄存器中计算
# 3. 消除了32MB中间张量的分配
# 4. 内存带宽需求减半(只读x和values,不写中间结果)

M.2 torch._check 的作用

torch._check(x.shape[0] >= 1)       # 至少1行
torch._check(x.shape[0] == values.shape[0])  # 行数匹配
  • torch._check 是编译时断言,帮助 torch.compile 生成更优化的代码
  • 如果编译器能证明条件总是成立,可以消除边界检查
  • 如果运行时违反,抛出RuntimeError

M.3 在Sampler中的使用位置

batched_count_greater_than 在 compute_token_logprobs(位于 v1/worker/gpu/sample/logprob.py)中使用,用于计算每个采样token在概率分布中的排名。

排名 = V - count_greater_than(logprobs, token_logprob) + 1

即:统计比该token logprob更大的token数,然后用V减去得到排名。


附录X Sampler.apply_logits_processors() 完整伪代码追踪

X.1 完整执行路径(含所有分支判断)

FUNCTION apply_logits_processors(logits, sampling_metadata, predict_bonus_token)

  # ===== Step 3: Allowed Token IDs白名单 =====
  IF sampling_metadata.allowed_token_ids_mask IS NOT None:
    # allowed_token_ids_mask: [max_batch_size, vocab_size]
    # True位置 = "该token不被允许"
    # 原地操作,不分配新内存
    logits.masked_fill_(sampling_metadata.allowed_token_ids_mask, float("-inf"))
    # 效果: 不在白名单中的token概率归零
    # 数学: softmax后,-inf位置概率为0
  
  # ===== Step 4: Bad Words排除 =====
  IF sampling_metadata.bad_words_token_ids IS NOT EMPTY:
    # bad_words_token_ids: {req_index: [[token_seq_1], [token_seq_2], ...]}
    # 对每个有禁词的请求:
    #   检查已生成序列尾部是否匹配禁词前缀
    #   匹配 → 将禁词最后一个token的logit设为-inf
    apply_bad_words(
      logits,                                    # [B, V]
      sampling_metadata.bad_words_token_ids,     # {idx: [[...], ...]}
      sampling_metadata.output_token_ids,        # [[...], [...], ...]
    )
    # 注意: apply_bad_words是原地修改,不返回新张量
  
  # ===== Step 5: Non-Argmax-Invariant Logits处理器 =====
  # 这些处理器可能改变argmax(贪心采样结果)
  # 必须在贪心采样之前应用
  FOR processor IN sampling_metadata.logitsprocs.non_argmax_invariant:
    logits = processor.apply(logits)
    # processor.apply() 可能返回新张量或原地修改
    # 必须用返回值更新logits引用
  
  # 内置non-argmax-invariant处理器:
  # 1. MinTokensLogitsProcessor:
  #    - 对未达到min_tokens的请求屏蔽stop token
  #    - logits.index_put_((req_indices, stop_token_ids), -inf)
  #    - 影响argmax: 如果原本argmax指向EOS,现在会被跳过
  
  # 2. LogitBiasLogitsProcessor:
  #    - 对特定token添加偏置值
  #    - logits[req_indices, token_indices] += bias_values
  #    - 影响argmax: 偏置可能使某个token的logit变为最大
  
  # ===== Thinking Budget =====
  holder = sampling_metadata.thinking_budget_state_holder
  IF holder IS NOT None AND holder.has_tracked_requests():
    # 有请求设置了thinking_token_budget
    # 且当前批次中有正在追踪的请求
    logits = holder.apply_to_logits(
      logits,
      predict_bonus_token=predict_bonus_token,
      spec_token_ids=sampling_metadata.spec_token_ids,
    )
    # holder.apply_to_logits() 内部:
    #   1. 清空mask和force_token_ids
    #   2. 计算每个请求在展平logits中的位置偏移
    #   3. 对in_end状态的请求,设置mask=True和force_token_ids
    #   4. logits[active_indices, force_tokens] = 1e9
    # 效果: 需要强制结束思考的token获得极大logit,确保被采样
  
  # ===== Step 6: 惩罚项 =====
  logits = self.apply_penalties(logits, sampling_metadata)
  # apply_penalties内部:
  #   IF no_penalties → return logits (快速路径)
  #   ELSE → apply_all_penalties(logits, prompt_ids, penalties, output_ids)
  #     内部: _convert_to_tensors → masked_fill_(-1, vocab_size) → apply_penalties
  #     底层: CUDA kernel并行计算三种惩罚
  
  RETURN logits
END FUNCTION

X.2 各步骤的数值影响分析

步骤操作数值范围变化概率影响
allowed_token_idsmasked_fill_(-inf)[a,b]→[a,-inf]白名单外token概率=0
bad_wordslogits[id]=-inf[a,b]→[a,-inf]特定token概率=0
MinTokensindex_put_(-inf)[a,b]→[a,-inf]stop token概率=0
LogitBiaslogits[slice]+=bias[a,b]→[a±c,b±d]改变token相对概率
ThinkingBudgetlogits[i,j]=1e9[a,b]→[1e9,b]强制token概率≈1
Repetitionlogit÷/×rep_p正÷,负×降低已出现token概率
Frequencylogit-=freq×count线性递减按次数降低概率
Presencelogit-=pres×I固定减量出现过就降低

附录Y Top-K/Top-P 过滤的数学完备性分析

Y.1 Top-K过滤后的概率归一化

Top-K过滤: logits[i] = -inf for i ∉ top-K

过滤后: probs = softmax(logits)
  probs[i] = exp(logits[i]) / sum(exp(logits))
  
  对于被屏蔽的token (logits = -inf):
    exp(-inf) = 0 → probs[i] = 0 ✓
  
  对于保留的token (logits > -inf):
    sum(exp(logits)) 只包含top-K个值
    → probs归一化到top-K个token上 ✓
  
  总概率: sum(probs) = 1 ✓ (因为softmax的归一化性质)

Y.2 Top-P过滤的累积概率保证

Top-P过滤: 保留累积概率≥1-p的token

定义: P_keep = sum(probs[i] for i in kept_tokens)
保证: P_keep ≥ p

证明:
  设排序后概率为 p_1 ≤ p_2 ≤ ... ≤ p_V
  保留条件: cumsum ≥ 1-p
  
  cumsum[j] = sum(p_1, ..., p_j)
  被屏蔽的token: cumsum[j] ≤ 1-p
  被保留的token: cumsum[j] > 1-p
  
  被屏蔽的概率总和: S_removed ≤ 1-p
  被保留的概率总和: S_keep = 1 - S_removed ≥ p ✓

Y.3 Top-K + Top-P的级联效应

先Top-K: 保留K个最大概率token
  → 概率分布变为 p'_1, ..., p'_K (归一化后)
  → 概率总和 = 1

再Top-P: 在Top-K的结果上应用
  → 累积概率从p'_1到p'_K
  → 保留累积概率≥1-p的子集
  
  注意: Top-P在Top-K之后应用
  因为: 如果先Top-P再Top-K,可能Top-P保留的token
        在Top-K中被移除,导致保留的token数少于预期
  
  例如: k=5, p=0.9
    先Top-K: 保留5个token,概率 [0.3, 0.25, 0.2, 0.15, 0.1]
    再Top-P: cumsum = [0.3, 0.55, 0.75, 0.9, 1.0]
    保留cumsum > 0.1的: [0.3, 0.55, 0.75, 0.9] → 保留4个token
    
    如果反过来:
    先Top-P: 保留cumsum > 0.1的 → 可能保留很多低概率token
    再Top-K: 只保留5个 → 可能移除Top-P想保留的token

Y.4 Min-P与Top-P的等价性分析

Min-P: 保留 prob[i] ≥ max_prob × min_p 的token
Top-P: 保留 cumsum(prob_sorted) > 1-p 的token

当min_p = 1 - p时,两者是否等价?
  答案: 不完全等价

  反例: probs = [0.5, 0.3, 0.15, 0.05]
  min_p = 0.1: threshold = 0.5 × 0.1 = 0.05
    保留: [0.5, 0.3, 0.15] (3个, 0.05被过滤)
  
  top_p = 0.9: cumsum = [0.05, 0.2, 0.5, 1.0]
    保留: [0.5, 0.3, 0.15] (3个, 0.05被过滤)
  
  看起来等价?但:
  probs = [0.5, 0.4, 0.05, 0.03, 0.02]
  min_p = 0.1: threshold = 0.5 × 0.1 = 0.05
    保留: [0.5, 0.4] (2个)
  
  top_p = 0.9: cumsum = [0.02, 0.05, 0.1, 0.5, 1.0]
    保留: [0.5, 0.4, 0.05] (3个, 0.9阈值在0.5和1.0之间)
  
  差异: Top-P保留了0.05,Min-P过滤了0.05
  原因: Min-P的阈值是绝对值(0.05),Top-P看的是累积概率

附录Z 采样管线的数值稳定性分析

Z.1 Float32精度需求

问题: 为什么强制使用float32而非float16/bfloat16?

1. Softmax溢出:
   exp(logits) 当logit > 88.7 (float32) → inf
   exp(logits) 当logit > 11.1 (float16) → inf
   vLLM的logits范围通常在[-20, 20],float16勉强可用
   但加上logit_bias后可能超过11.1

2. 温度除法:
   logits / temperature 当temperature很小(如0.01)时
   float16: 20 / 0.01 = 2000 > 65504 → inf
   float32: 20 / 0.01 = 2000 ✓

3. log_softmax:
   log(exp(x) / sum(exp(x)))
   需要log(sum(exp(x)))的精确计算
   float16的log精度不足

4. 惩罚计算:
   repetition_penalty: logit / rep_p 或 logit * rep_p
   当rep_p很大时,float16的精度损失导致惩罚效果不均匀

Z.2 Log-Sum-Exp技巧

log_softmax(x) = x - log(sum(exp(x)))

直接计算: log(sum(exp(x)))
  问题: exp(x)可能溢出

改进: log(sum(exp(x - max(x)))) + max(x)
  1. x' = x - max(x) → 所有值 ≤ 0
  2. exp(x') → 所有值 ∈ (0, 1]
  3. sum(exp(x')) → 不会溢出
  4. log(sum(exp(x'))) + max(x) → 正确的log_softmax

PyTorch的log_softmax已经内置了此技巧
vLLM直接使用: logits.log_softmax(dim=-1, dtype=torch.float32)

Z.3 Gumbel-Max的数值稳定性

probs.div_(q).argmax(dim=-1)

问题: 当q很小时,probs/q可能溢出

解决:
  q = exponential_(shape)  # Exp(1)分布
  q的最小值 ≈ 1e-38 (float32)
  probs最大值 ≈ 1.0
  probs/q最大值 ≈ 1e38 → 接近float32上限

  但argmax不关心绝对值,只关心相对大小
  即使某些值溢出为inf,argmax仍返回正确的索引
  因为: 如果probs[i]/q[i] = inf,说明q[i]极小
  → 该token被采样的概率极高 → 应该被选中

Z.4 Temperature=0的边界处理

问题: logits / 0 = inf → softmax全为NaN

vLLM的处理:
  1. 检测all_greedy → 直接argmax,不做温度除法
  2. 混合模式: expand时将0替换为1
     expand_batch_to_tokens(temperature, ..., replace_from=0, replace_to=1)
  3. 采样后: torch.where(temperature >= ε, random_sampled, greedy_sampled)
     温度<ε的请求使用贪心结果,温度≥ε的使用随机结果

  _SAMPLING_EPS = 1e-5
  # 温度在[0, 1e-5)范围内视为贪心
  # 避免极小温度导致的数值问题

附录AB SamplingMetadata 构建流程深度追踪

AB.1 从Scheduler到Sampler的数据流

Scheduler._schedule() 
  → 创建 SchedulerRunningOutputs
  → 包含每个运行中请求的采样参数

GPUModelRunner._prepare_inputs()
  → 遍历运行中请求
  → 收集 sampling_params → 构建 SamplingMetadata

SamplingMetadata 构造器:
  1. 收集温度: [temp_0, temp_1, ..., temp_B-1]
  2. 收集top-p: [p_0, p_1, ..., p_B-1]
  3. 收集top-k: [k_0, k_1, ..., k_B-1]
  4. 收集惩罚: presence/frequency/repetition penalties
  5. 收集logprobs需求: max_num_logprobs
  6. 收集bad_words: {req_idx: [[tok_seq], ...]}
  7. 收集allowed_token_ids: {req_idx: set(int)}
  8. 收集seed: {req_idx: int}
  9. 构建output_token_ids引用列表
  10. 构建prompt_token_ids张量
  11. 调用 build_logitsprocs() 构建处理器容器
  12. 调用 maybe_create_thinking_budget_state_holder()

AB.2 _convert_to_tensors() 惩罚张量转换

输入 (Python列表):
  presence_penalties = [0.0, 0.5, 0.0, 0.3]
  frequency_penalties = [0.2, 0.0, 0.1, 0.0]
  repetition_penalties = [1.0, 1.2, 1.0, 1.1]

转换步骤:
  1. torch.tensor(penalties, dtype=torch.float32, device="cpu")
  2. .to(device=device, non_blocking=True)
  3. 结果: [B] 形状的GPU tensor

  特殊处理:
  - 如果所有值相同(如repetition全为1.0): 创建标量tensor而非向量
  - 如果列表为空: 创建空tensor (shape=[0])

AB.3 allowed_token_ids_mask 构建流程

输入: allowed_token_ids = {0: {100, 200, 300}, 2: {50, 60}}

步骤1: 确定mask形状
  max_num_reqs = max(allowed_token_ids.keys()) + 1 = 3
  vocab_size = 32000

步骤2: 创建反向mask(True=不被允许)
  mask = torch.zeros(3, 32000, dtype=torch.bool)
  mask[0, :] = True  # 请求0: 默认全部不允许
  mask[0, [100, 200, 300]] = False  # 请求0: 100/200/300允许
  mask[1, :] = False  # 请求1: 无限制(不在allowed_token_ids中)
  mask[2, :] = True  # 请求2: 默认全部不允许
  mask[2, [50, 60]] = False  # 请求2: 50/60允许

步骤3: 应用
  logits.masked_fill_(mask, float("-inf"))
  # 不允许的token → logit = -inf → softmax概率 = 0

AB.4 Seed到Generator的映射

输入: seeds = {0: 42, 3: 123}

步骤1: 为每个有seed的请求创建CPU Generator
  generators = {
    0: torch.Generator(device="cpu").manual_seed(42),
    3: torch.Generator(device="cpu").manual_seed(123),
  }

步骤2: 在采样时使用
  # Gumbel-Max采样
  q = torch.empty([B, V], dtype=torch.float32)
  for i, generator in generators.items():
    q[i].exponential_(generator=generator)
  # 有seed的请求使用确定性随机数
  # 无seed的请求使用全局随机数(不可复现)

附录AC apply_bad_words() 实现深度分析

AC.1 单token禁词的快速路径

# 当bad_word只有一个token时:
# bad_word = [token_id]
# prefix_length = 0
# 只需将 logits[row, token_id] = -inf
# 这是O(1)操作

if prefix_length == 0:
    # 单token禁词 → 直接屏蔽
    logits[row, last_token_id] = float("-inf")
    continue  # 处理下一个禁词

AC.2 多token禁词的前缀匹配

# 当bad_word有N个token时:
# bad_word = [t1, t2, ..., tN]
# 需要检查output_token_ids尾部是否匹配[t1, ..., tN-1]
# 如果匹配,屏蔽tN

prefix = bad_word[:-1]  # 前缀(不含最后一个token)
last_token = bad_word[-1]  # 需要屏蔽的token

# 获取已生成序列的尾部
tail = output_token_ids[-len(prefix):]

if tail == prefix:
    # 匹配!下一个token如果等于last_token则构成禁词
    # 屏蔽last_token
    logits[row, last_token] = float("-inf")

AC.3 多禁词的并行检查

请求0有3个禁词:
  bad_words[0] = [[5], [10, 20, 99], [30, 40]]

检查流程:
  word 0: [5] → prefix=[], last=5 → 直接屏蔽5
  word 1: [10, 20, 99] → prefix=[10,20], last=99
    tail = output[-2:] 
    如果 [10, 20] == [10, 20] → 屏蔽99
  word 2: [30, 40] → prefix=[30], last=40
    tail = output[-1:]
    如果 [30] == [30] → 屏蔽40

所有匹配的禁词token同时被屏蔽
如果同一token被多个禁词指向,只屏蔽一次(-inf ∪ -inf = -inf)

附录AD unconditional_to_conditional_rates 数学推导

AD.1 转换公式

设无条件接受率为 α₁, α₂, …, αₖ
求条件接受率 β₁, β₂, …, βₖ

使得联合接受概率满足:

P(accept all K) = Π_{i=1}^{K} β_i

而原始设计:
P(accept all K) = Π_{i=1}^{K} α_i

条件接受率的定义:

β₁ = α₁  (第1个位置无条件)
β₂ = P(accept pos2 | accept pos1) × β₁

但更准确的定义是"在前面已接受的条件下的条件接受率":

P(accept pos1) = α₁ → β₁ = α₁

P(accept pos1 AND accept pos2) = α₁ × α₂ (独立)
P(accept pos2 | accept pos1) = α₂ (因为独立)

但这里"条件"的含义不同:
β₂ 是"在pos1已接受的前提下,pos2的条件接受率"
= P(pos2 accepted | pos1 accepted)
= α₂  (因为pos1和pos2独立)

Wait, 但 unconditional_to_conditional_rates 的实际实现是:

conditional_rates[0] = unconditional_rates[0]
conditional_rates[i] = unconditional_rates[i] / (1 - sum of previous conditional rates that were not met)

实际代码中的逻辑更复杂,因为它考虑的是"在已经到达位置i的前提下接受的概率"。

简化理解:
  如果前i-1个都接受了,则到达位置i的概率是:
  Π_{j=1}^{i-1} β_j
  
  位置i被接受的总概率:
  α_i = Π_{j=1}^{i-1} β_j × β_i
  
  因此:
  β_i = α_i / Π_{j=1}^{i-1} β_j

AD.2 数值示例

unconditional_rates = [0.8, 0.6, 0.4]

β₁ = 0.8

β₂ = 0.6 / 0.8 = 0.75
  # 解释: "到达位置2的概率是0.8(位置1接受)"
  # "在到达位置2的前提下,接受概率0.75"
  # "位置2被接受的总概率: 0.8 × 0.75 = 0.6 = α₂ ✓"

β₃ = 0.4 / (0.8 × 0.75) = 0.4 / 0.6 ≈ 0.667
  # "到达位置3的概率: 0.8 × 0.75 = 0.6"
  # "在到达位置3的前提下,接受概率0.667"
  # "位置3被接受的总概率: 0.6 × 0.667 = 0.4 = α₃ ✓"

验证联合概率:
  P(accept all 3) = 0.8 × 0.75 × 0.667 = 0.4
  P(accept all 3) = α₁ × α₂ × α₃ = 0.8 × 0.6 × 0.4 = 0.192
  
  ⚠️ 这两个不相等!
  
  实际上无条件接受率的含义是:
  α_i = P(position i is accepted)
  NOT P(position i is accepted AND all previous are accepted)
  
  所以联合接受概率 ≠ Π α_i
  联合接受概率 = Π β_i = 0.8 × 0.75 × 0.667 = 0.4
  
  而每个位置被接受的概率:
  P(accept pos1) = 0.8 = α₁ ✓
  P(accept pos2) = 0.8 × 0.75 = 0.6 = α₂ ✓  
  P(accept pos3) = 0.8 × 0.75 × 0.667 = 0.4 = α₃ ✓

附录AE 全局术语表与中英对照

英文术语中文翻译代码位置
Sampling采样sampler.py
Logits对数概率全局
Greedy Sampling贪心采样sampler.py
Random Sampling随机采样sampler.py
Top-K SamplingTop-K采样topk_topp_sampler.py
Top-P (Nucleus) SamplingTop-P核采样topk_topp_sampler.py
Min-P SamplingMin-P采样builtin.py
Temperature温度sampler.py
Repetition Penalty重复惩罚penalties.py
Frequency Penalty频率惩罚penalties.py
Presence Penalty存在惩罚penalties.py
Logit Bias对数偏置builtin.py
Min Tokens最小令牌数builtin.py
Bad Words禁词sampler.py
Allowed Token IDs白名单令牌sampler.py
Speculative Decoding投机解码rejection_sampler.py
Draft Model草稿模型rejection_sampler.py
Target Model目标模型rejection_sampler.py
Rejection Sampling拒绝采样rejection_sampler.py
Recovered Token恢复令牌rejection_sampler.py
Bonus Token额外令牌rejection_sampler.py
Thinking Budget思考预算thinking_budget_state.py
Gumbel-Max TrickGumbel-Max技巧topk_topp_sampler.py
Ternary Search三分搜索topk_topp_triton.py
Pivot枢轴topk_topp_triton.py
Outlier离群值topk_topp_triton.py
Persistent Batch持久批次interface.py
BatchUpdate批次更新interface.py
argmax-invariantargmax不变interface.py
LogitsProcessorLogits处理器interface.py
AdapterLogitsProcessor适配器Logits处理器init.py
FQCN完全限定类名init.py
Entry Points入口点init.py
pin_memory锁页内存builtin.py
Triton KernelTriton内核topk_topp_triton.py
CSR FormatCSR格式rejection_sampler.py
Cumulative Offset累积偏移rejection_sampler.py
PLACEHOLDER_TOKEN_ID占位符令牌ID(-1)rejection_sampler.py
Synthetic Mode合成模式rejection_sampler.py
Conditional Rate条件接受率rejection_sampler.py
QritaQrita算法topk_topp_triton.py
Buffer Cache缓冲区缓存topk_topp_triton.py
Table Cache查找表缓存topk_topp_triton.py
Log-Sum-Exp对数求和指数sampler.py
Float64 Precision双精度rejection_sampler.py
Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐