【vllm】(v1 Sample)vLLM V1 Sample—Part 1架构总览与核心采样器
vLLM V1 Sample 模块超深度架构分析 — Part 1: 架构总览与核心采样器
分析范围:
vllm/v1/sample/目录全部源码(15个Python文件,约4,500行)
分析日期: 2026-05-25
目录
- 第一章 模块定位与全局架构
- 第二章 核心数据结构深度解析
- 第三章 Sampler 核心采样器逐行解析
- 第四章 TopKTopPSampler 采样算子深度解析
- 4.1 多平台策略模式
- 4.2 forward_native() PyTorch原生实现
- 4.3 forward_cuda() FlashInfer加速路径
- 4.4 forward_cpu() CPU路径
- 4.5 forward_hip() ROCm/aiter路径
- 4.6 forward_xpu() Intel XPU路径
- 4.7 apply_top_k_top_p() 调度函数
- 4.8 apply_top_k_top_p_pytorch() 排序实现
- 4.9 apply_top_k_only() 无排序优化
- 4.10 random_sample() Gumbel-Max采样
- 4.11 flashinfer_sample() FlashInfer采样
- 4.12 compiled_random_sample() 编译优化
- 第五章 Penalties惩罚算子深度解析
- 第六章 BadWords与Logprobs辅助算子
- 附录A 采样管线完整时序图
- 附录B 术语表
第一章 模块定位与全局架构
1.1 业务职责与功能定位
vllm/v1/sample 模块是 vLLM V1 架构中推理采样阶段的核心实现。其业务职责为:
- Logits后处理:对模型输出的原始logits施加一系列变换(logits处理器、惩罚项、温度缩放、top-k/top-p过滤)
- Token采样决策:从处理后的概率分布中采样下一个token(贪心采样或随机采样)
- 日志概率计算:计算采样token及top-N token的logprobs
- 投机解码采样:实现基于拒绝采样的推测解码(speculative decoding)验证与修正
- 思考预算控制:管理推理模型的thinking token预算,强制结束思考阶段
功能定位一句话总结:sample模块是从"模型输出logits"到"最终采样token"之间的完整决策管线。
1.2 在系统中的位置
上游依赖:
GPUModelRunner.execute_model()→ 调用Sampler.forward(),传入模型输出的logits和SamplingMetadataSamplingMetadata由GPUInputBatch在每步构建
下游输出:
SamplerOutput(包含sampled_token_ids+logprobs_tensors)→ 返回给GPUModelRunnerRejectionSampler.forward()→ 投机解码路径的替代采样入口
1.3 模块全景架构图
1.4 文件结构与代码量统计
| 文件路径 | 行数 | 核心类/函数 | 职责 |
|---|---|---|---|
sampler.py | 425 | Sampler | 主采样器入口,协调所有采样步骤 |
metadata.py | 55 | SamplingMetadata | 采样参数数据容器 |
ops/topk_topp_sampler.py | 458 | TopKTopPSampler + 辅助函数 | top-k/top-p过滤+随机采样 |
ops/topk_topp_triton.py | 1,058 | _topk_topp_kernel + apply_top_k_top_p_triton | Triton加速的top-k/top-p |
ops/penalties.py | 57 | apply_all_penalties | 惩罚项应用 |
ops/bad_words.py | 57 | apply_bad_words / apply_bad_words_with_drafts | 禁词屏蔽 |
ops/logprobs.py | 27 | batched_count_greater_than | logprobs辅助计算 |
rejection_sampler.py | 921 | RejectionSampler + Triton kernels | 投机解码拒绝采样 |
thinking_budget_state.py | 528 | ThinkingBudgetStateHolder | 思考token预算管理 |
logits_processor/__init__.py | 357 | build_logitsprocs / AdapterLogitsProcessor | logits处理器构建与适配器 |
logits_processor/interface.py | 106 | LogitsProcessor / BatchUpdate | 处理器抽象接口 |
logits_processor/state.py | 165 | BatchUpdateBuilder / LogitsProcessors | 批次更新构建器 |
logits_processor/builtin.py | 332 | MinP / 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, ...]:稀疏存储,大多数请求没有禁词约束
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 数据流向全景图
第三章 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
"""
__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.Modulelogprobs_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 在惩罚/温度之前计算。这样做的原因是:
- 用户期望看到的是"模型对token的原始评价",而非经过人为调整后的值
- 与OpenAI API的logprobs语义对齐
- 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
)
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)
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影响 |
|---|---|---|---|
| Repetition | logit > 0 → logit/rep_p; logit < 0 → logit*rep_p | 惩罚所有已出现token | 是(改变logit符号和大小) |
| Frequency | logit -= freq_p * count(token) | 按出现次数惩罚 | 是(减去正值) |
| Presence | logit -= 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
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 方法:
策略选择决策表:
| 平台 | 加速库 | 条件 | forward方法 |
|---|---|---|---|
| CUDA | FlashInfer | VLLM_USE_FLASHINFER_SAMPLER=1 + GPU capability ≥ 8.0 | forward_cuda |
| CUDA | PyTorch | 默认或FlashInfer不可用 | forward_native |
| CPU | torch.compile | x86/ARM架构 | forward_cpu |
| CPU | PyTorch | RISC-V/PowerPC | forward_native |
| ROCm | aiter | aiter_ops已安装 | forward_hip |
| XPU | 自定义kernel | VLLM_XPU_USE_SAMPLER_KERNEL=1 | forward_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 的性能差异:
| 特性 | FlashInfer | PyTorch-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)
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采样
优势:
- 全GPU操作,无CPU同步
- 可批量处理
- 可复现(通过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 原理:
- 不对词表排序,而是生成均匀随机数u
- 对每个token,以概率p_i接受(p_i > u的阈值)
- 如果恰好一个token被接受则选中
- 否则重试
- 复杂度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 # 只惩罚一次
# 效果:只要出现过就惩罚固定量,不关心出现次数
第六章 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序列(如"不"+“好”)。屏蔽逻辑是:
- 检查已生成序列的尾部是否匹配禁词的前缀
- 如果匹配,将禁词的最后一个token的logit设为-inf
- 这样模型就不会生成该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 采样管线完整时序图
附录B 术语表
| 术语 | 全称 | 含义 |
|---|---|---|
| Logits | — | 模型输出的未归一化对数概率 |
| Logprobs | Log Probabilities | log(softmax(logits)),归一化对数概率 |
| Top-K | — | 只保留概率最高的K个token |
| Top-P (Nucleus) | — | 保留概率累积和≥P的最少token |
| Min-P | — | 过滤概率低于 max_prob×min_p 的token |
| Gumbel-Max | Gumbel-Max Trick | 用指数分布噪声实现分类采样的技巧 |
| Repetition Penalty | — | 对已出现token施加惩罚,防重复 |
| Frequency Penalty | — | 按出现次数线性惩罚已出现token |
| Presence Penalty | — | 出现过就惩罚固定量 |
| Argmax-Invariant | — | logit处理器不改变argmax(贪心采样结果不变) |
| Pin Memory | Pinned Memory | 锁定CPU内存不被换出,加速CPU→GPU传输 |
| Spec Decode | Speculative 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_native | forward_cuda | forward_cpu | forward_hip | forward_xpu |
|---|---|---|---|---|---|
| 平台 | 通用 | NVIDIA CUDA | CPU | AMD ROCm | Intel XPU |
| 加速库 | PyTorch/Triton | FlashInfer | torch.compile | aiter | 自定义kernel |
| Top-K/Top-P实现 | apply_top_k_top_p | FlashInfer | apply_top_k_top_p_pytorch | aiter_sample | xpu_topk_topp_sampler |
| 采样方式 | Gumbel-Max | FlashInfer rejection | Gumbel-Max (compiled) | aiter sampling | 自定义 |
| Per-request seed | ✅ | ❌ | ✅ | ❌ | ❌ |
| Processed logprobs | ✅ | ❌ | ✅ | ❌ | ✅ |
| Batch阈值 | 无 | 无 | 无 | 无 | 无 |
| 回退条件 | — | k=None & p=None / generators | generators | k=None & p=None / generators / DISABLE | generators |
| 算子复杂度 | O(V log V) 或 O(V) Triton | O(V) rejection | O(V log V) 或 O(V×k) | O(V) | O(V) |
| 数值精度 | float32 | float32 | float32 | float32 | float32 |
J.1 各平台采样算法对比
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%
附录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
附录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_ids | masked_fill_(-inf) | [a,b]→[a,-inf] | 白名单外token概率=0 |
| bad_words | logits[id]=-inf | [a,b]→[a,-inf] | 特定token概率=0 |
| MinTokens | index_put_(-inf) | [a,b]→[a,-inf] | stop token概率=0 |
| LogitBias | logits[slice]+=bias | [a,b]→[a±c,b±d] | 改变token相对概率 |
| ThinkingBudget | logits[i,j]=1e9 | [a,b]→[1e9,b] | 强制token概率≈1 |
| Repetition | logit÷/×rep_p | 正÷,负× | 降低已出现token概率 |
| Frequency | logit-=freq×count | 线性递减 | 按次数降低概率 |
| Presence | logit-=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 Sampling | Top-K采样 | topk_topp_sampler.py |
| Top-P (Nucleus) Sampling | Top-P核采样 | topk_topp_sampler.py |
| Min-P Sampling | Min-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 Trick | Gumbel-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-invariant | argmax不变 | interface.py |
| LogitsProcessor | Logits处理器 | interface.py |
| AdapterLogitsProcessor | 适配器Logits处理器 | init.py |
| FQCN | 完全限定类名 | init.py |
| Entry Points | 入口点 | init.py |
| pin_memory | 锁页内存 | builtin.py |
| Triton Kernel | Triton内核 | topk_topp_triton.py |
| CSR Format | CSR格式 | 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 |
| Qrita | Qrita算法 | 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 |
更多推荐
所有评论(0)