资讯动态

xPress:并行细化加速扩散草稿模型,突破投机解码延迟瓶颈

发布时间:2026/8/30 16:40:07 来源:尧图企业网站定制
之前一直在做大模型推理性能优化最头疼的就是自回归解码的串行瓶颈每个 token 都要等前一个 token 算完才能继续GPU 利用率上不去业务延迟又压不下来。后来接触到 Speculative Decoding思路很惊艳——用小模型先草拟一批 token再让大模型一次性验证把串行推理变成“草稿 验证”的流水线。但真正落到工程上又会遇到另一个问题当草稿模型使用扩散模型Diffusion Drafter时草稿生成本身需要多步去噪这一步反而成为新的时间开销。本文要讲的 xPress核心就是用 Parallel Refinement 的思路来压低扩散草稿生成延迟让 Speculative Decoding 在扩散草稿模型上也能跑出真实加速效果。这篇文章适合正在做 LLM 推理加速的算法工程师、对扩散模型生成离散 token 感兴趣的研究者以及想自己复现投机解码流程的开发者。读完你会理解扩散草稿模型为什么慢、xPress 的并行细化到底并行在哪里、如何写一个最小化的并行细化验证程序以及工程落地时最容易踩的坑。1. 什么是 Speculative Decoding为什么要关注 Diffusion Drafter1.1 自回归推理的瓶颈大语言模型LLM生成文本时通常是自回归方式每一步只生成一个 token然后把这个 token 拼接到输入序列末尾再计算下一步。这个过程的计算特征非常“稀疏”——GPU 每执行一次 forward只为得到一个 token 的概率分布虽然 batch 内可能有很多序列但每条序列在时间轴上无法并行推进。举个例子目标生成 128 个 token就意味着要做 128 次模型前向计算。如果模型是 70B 参数即使开启 bf16 和张量并行单次前向的显存占用和耗时也相当可观。于是工程上出现了很多优化手段KV Cache、连续批处理、量化、投机解码。其中投机解码是近年来最接近“无损加速”的方案之一。1.2 Speculative Decoding 的基本思想Speculative Decoding 的思路很直接用一个小而快的草稿模型Drafter先生成一段候选 token 序列长度记为 k然后用大模型Target Model对这段候选序列做一次并行前向计算逐 token 验证草稿结果是否与目标模型分布一致。验证时采用拒绝采样或局部采样策略。如果某个位置草稿概率与目标概率偏差过大就截断接受并用目标模型在该位置的采样结果作为最终输出同时丢弃后续候选 token。这样每一轮大模型虽然只 forward 一次却能“附带”产出多个有效 token实际生成速度上限取决于草稿被接受的比例也就是接受率Acceptance Rate。在这个流程里草稿生成是决定收益上限的关键组件。传统草稿模型往往就是一个相同 tokenizer 的小型自回归模型生成草稿仍然需要串行 k 次前向只是小模型单次计算更快。于是研究者开始探索能不能用非自回归模型一次性生成整段 draft实现 O(1) 级别的草稿延迟Diffusion Drafter 正是从这个角度切入的。1.3 Diffusion Drafter 的特殊之处扩散模型原本擅长图像生成通过多步去噪把高斯噪声逐步还原成数据分布。离散 token 序列也可以被嵌入到一个连续空间中经过若干次去噪后在每一步对 token 空间做 softmax 映射从而得到一组候选 token。这样设计的好处是理论上草稿生成不再是 k 次串行而是多次去噪迭代并行作用在整段序列上坏处是扩散模型的“多步去噪”本身不是免费的通常需要 4 步、8 步甚至更多步迭代。如果每一步都在同一个 GPU 上串行执行那草稿生成延迟未必比小自回归模型低。xPress 这篇工作关注的就是这个矛盾点。它提出的 Parallel Refinement 不是把多个草稿拆到不同机器上而是把“多轮去噪”中的每一轮同时对多个候选草稿并行推进从而减少扩散草稿生成在端到端延迟中的占比。2. xPress 的核心Parallel Refinement 的原理推演2.1 为什么扩散草稿生成存在“细化延迟”扩散模型生成草稿时通常先随机初始化一段长度为 k 的连续向量然后逐步去噪每一轮都会让向量更接近真实 token embedding 的分布。每一步本质上都在“细化”当前草稿。如果只有一份草稿那么整个流程是初始化草稿 → 去噪第 1 步 → 去噪第 2 步 → …… → 去噪第 n 步 → 输出 token 候选这是一条串行链路。n 步去噪意味着 n 次模型 forward。相比自回归草稿模型生成 k 个 token 的 k 次 forward扩散模型确实把次数从 k 降到了 n但 n 的绝对延迟依然会占据每个 decoding step 的很大比例。更重要的是在 Speculative Decoding 的循环中草稿阶段和验证阶段是有依赖关系的必须先拿到草稿目标模型才能验证。因此草稿生成耗时越长整个系统的端到端延迟就越差。如果草稿生成从 1 次 forward 变成 8 次 forward即使每次 forward 比目标模型小很多也会让整体吞吐大幅下降。2.2 从串行细化到并行细化xPress 的思路是让“细化”这个动作本身并行化。具体来说如果我们在每个 decoding step 同时维护多个候选草稿例如 4 个草稿序列那么每个草稿都需要去噪 n 步。常规做法是一个草稿完整走完 n 步再去处理下一个草稿这是“串行细化”更好的做法是让所有草稿在同一时刻执行同样的去噪步例如第 1 轮草稿 0/1/2/3 同时去噪第 1 步 第 2 轮草稿 0/1/2/3 同时去噪第 2 步 …… 第 n 轮草稿 0/1/2/3 同时去噪第 n 步每一轮都相当于把一个 batch 喂给扩散模型利用 GPU 的并行计算能力同时处理多条草稿。这样总延迟不再是“单条草稿的 n 步 × 草稿数”而是接近“单条草稿的 n 步 batch 通信开销”。这就是 Parallel Refinement 的关键并行发生在草稿之间而非某个草稿内部的 token 之间。后者受限于扩散模型自身的序列建模依赖很难完全并行前者则非常自然因为多条草稿互不依赖。2.3 并行细化的验证与接受机制有了多条候选草稿之后验证阶段也需要调整。传统投机解码只验证一条草稿现在需要把多个草稿拼成一个更大的 batch交给目标模型并行打分。目标模型会计算所有草稿位置上的 token 概率然后从多个候选中选择最佳前缀。一种合理的策略是对每条草稿分别计算接受长度优先接受接受长度最长的草稿如果两条草稿接受长度相同则按照目标模型在该位置的概率采样。这样目标模型的单次 forward 同时完成了“验证 选择”不会比验证单条草稿多出额外次数只是 batch 变大了。从实现角度看这其实是一种用显存换延迟的方案。并行细化增加了每次 forward 的输入 token 数量但减少了草稿生成阶段的串行等待在 GPU 资源相对充裕、 batch size 余量较大的场景下收益明显。3. 环境准备与版本说明3.1 运行环境复现 xPress 或做类似实验时建议准备以下环境。版本不需要完全照抄但要保证大版本兼容。操作系统LinuxUbuntu 20.04 / 22.04或 CentOS 7 GPUNVIDIA Ampere 架构或更新例如 A100、A800、H100 CUDA11.8 或 12.x Python3.9 / 3.10 / 3.11 PyTorch2.x Transformers4.30 或更高如果只是验证本文的调度模拟代码没有 GPU 也可以运行因为示例里只涉及 Python 多线程和标准库计时。3.2 依赖与参考实现真实复现 xPress 时需要依赖 HuggingFace Transformers、扩散模型相关组件例如 diffusers 或自研的 discrete diffusion 模块。由于不同模型的 API 差异较大这里不写死具体版本建议以论文官方仓库和本地环境为准。pip install torch transformers accelerate如果你的目标模型较大例如 7B 或 70B 级别还需要安装bitsandbytes做量化加载或者使用vLLM等推理框架。3.3 项目目录结构建议按下面的结构组织实验代码xpress_demo/ ├── README.md ├── configs/ │ └── default.yaml ├── models/ │ ├── drafter.py │ └── target_wrapper.py ├── decoding/ │ ├── speculative_decode.py │ ├── parallel_refinement.py │ └── verify.py ├── experiments/ │ └── run_benchmark.py └── scripts/ └── download_models.sh这样把模型封装、解码逻辑、验证逻辑分开后续做性能 profiling 会方便很多。4. 动手复现最小化的并行细化示例4.1 示例 1Speculative Decoding 主干流程先看一个简化版的投机解码主循环。下面的代码不绑定具体模型接口重点表达“草稿生成 → 目标模型验证 → 接受前 k 个 token”这个流程。import torch import torch.nn.functional as F def speculative_decode_loop( prompt_tokens, draft_model, target_model, max_new_tokens32, draft_len4, temperature1.0, ): 简化版投机解码主循环。 参数 prompt_tokens: list[int]输入 token draft_model: 草稿模型接收最近一组 token返回下一个 token logits target_model: 目标模型接收完整序列返回每个位置的 logits max_new_tokens: 最大新生成 token 数 draft_len: 每轮生成的草稿长度 k temperature: 采样温度 generated list(prompt_tokens) while len(generated) - len(prompt_tokens) max_new_tokens: # 1. 草稿阶段自回归生成 draft_len 个候选 token draft_tokens [] draft_logprobs [] # 这里为了演示保存每个位置的草稿 logits local_input [generated[-1]] for _ in range(draft_len): with torch.no_grad(): logits draft_model(torch.tensor([local_input]))[0, -1] logprobs F.log_softmax(logits / temperature, dim-1) next_token torch.multinomial(logprobs.exp(), 1).item() draft_tokens.append(next_token) draft_logprobs.append(logprobs[next_token].item()) local_input.append(next_token) # 2. 验证阶段目标模型对整条序列做一次前向 candidate_input generated draft_tokens with torch.no_grad(): target_logits target_model(torch.tensor([candidate_input]))[0] # 3. 逐 token 接受决策这里使用最简单的概率阈值方式 accept_len 0 for i in range(draft_len): t_logprobs F.log_softmax( target_logits[-(draft_len - i)] / temperature, dim-1 ) t_prob t_logprobs[draft_tokens[i]].exp().item() d_prob draft_logprobs[i] # 草稿概率大于目标概率时按比例随机接受否则接受概率为 1 if d_prob t_prob: accept_len 1 else: r torch.rand(1).item() if r (t_prob / (d_prob 1e-9)): accept_len 1 else: break generated.extend(draft_tokens[:accept_len]) # 如果没有完全接受补一个目标模型采样的 token if accept_len draft_len: with torch.no_grad(): repair_logits target_model( torch.tensor([generated]) )[0, -1] repair_token torch.multinomial( F.softmax(repair_logits / temperature, dim-1), 1 ).item() generated.append(repair_token) return generated这段代码省略了 EOS 判断和 KV Cache 优化目的是把投机解码的骨架讲清楚。实际项目中你要把draft_model和target_model替换成改造后的模型接口并为目标模型维护 KV Cache避免每次都从头计算整条序列。4.2 示例 2扩散草稿模型的去噪循环扩散式草稿生成器与自回归草稿模型最大的不同是它会先初始化一组连续草稿向量再通过多次去噪得到 logits。下面是一个结构示意import torch import torch.nn.functional as F def diffusion_draft( prompt_embeds, diffusion_model, token_head, num_denoise_steps8, draft_len4, temperature1.0, ): 扩散草稿生成示意。 参数 prompt_embeds: Tensorshape 为 (batch, seq_len, hidden) diffusion_model: 把噪声向量逐步还原到 token embedding 空间 token_head: 从 hidden 向量映射到 vocab logits num_denoise_steps: 去噪步数 draft_len: 草稿长度 # 1. 初始化噪声草稿向量 batch prompt_embeds.shape[0] x torch.randn(batch, draft_len, prompt_embeds.shape[-1]) # 2. 多步去噪每一轮都在细化整段草稿 for step in range(num_denoise_steps): # 在真实实现中这里需要把 prompt 信息和当前时间步一起喂给模型 t torch.full((batch,), step) denoised diffusion_model(prompt_embeds, x, t) # 去噪更新具体公式取决于模型设计 x x 0.5 * (denoised - x) # 此时可以对 x 做一次临时 logits 映射用于监控草稿变化 if step % 2 0: logits token_head(x) probs F.softmax(logits / temperature, dim-1) print(fstep {step}: top token {probs.argmax(-1).tolist()}) # 3. 最后一次去噪后映射到 vocab logits得到候选 token logits token_head(x) probs F.softmax(logits / temperature, dim-1) draft_tokens probs.argmax(dim-1) draft_logprobs torch.log(probs.gather(-1, draft_tokens.unsqueeze(-1)).squeeze(-1)) return draft_tokens, draft_logprobs这个示例的关键点在于循环体内部的diffusion_model(prompt_embeds, x, t)是并行作用于整段x的没有逐 token 依赖。所以在num_denoise_steps8的场景下草稿延迟理论上是 8 次 forward而不是draft_len4的 4 次串行 forward。这也是扩散草稿模型最大的结构优势。4.3 示例 3串行细化与并行细化的性能对比前面说过xPress 的 Parallel Refinement 是把多条草稿放一起在同一去噪步并行推进。下面用一个小程序模拟这个调度差异。import time from concurrent.futures import ThreadPoolExecutor class Draft: def __init__(self, draft_id, noise_level10): self.draft_id draft_id self.noise_level noise_level def denoise_step(self, step): # 模拟一次扩散去噪 forward实际应用中这里是 GPU 计算 time.sleep(0.05) self.noise_level - 1 return self def sequential_refine(drafts, denoise_steps8): 串行细化每份草稿完整走完所有去噪步再处理下一份。 for draft in drafts: for step in range(denoise_steps): draft.denoise_step(step) return drafts def parallel_refine(drafts, denoise_steps8): 并行细化每一轮同时推进所有草稿的同一去噪步。 with ThreadPoolExecutor(max_workerslen(drafts)) as executor: for step in range(denoise_steps): futures [ executor.submit(draft.denoise_step, step) for draft in drafts ] for future in futures: future.result() return drafts def run_benchmark(draft_count4, denoise_steps8): drafts [Draft(i) for i in range(draft_count)] start time.perf_counter() sequential_refine(drafts, denoise_steps) seq_time time.perf_counter() - start drafts [Draft(i) for i in range(draft_count)] start time.perf_counter() parallel_refine(drafts, denoise_steps) para_time time.perf_counter() - start return seq_time, para_time if __name__ __main__: seq_t, para_t run_benchmark(4, 8) print(fsequential refine time: {seq_t:.3f}s) print(fparallel refine time: {para_t:.3f}s) print(fspeedup: {seq_t / para_t:.2f}x)这段代码用time.sleep(0.05)模拟一次 forward真正的 GPU 场景会更复杂但调度思路完全一致。并行细化并不是把所有草稿的所有步都一次性塞给 GPU而是每一轮把当前步的多个草稿组成 batch利用 GPU 宽并行处理。4.4 运行与预期结果在 4 个草稿、8 步去噪的配置下串行细化大约需要4 × 8 × 0.05 1.6s并行细化如果单轮 4 个任务并行执行大约需要8 × 0.05 0.4s加上少量线程调度开销。sequential refine time: 1.602s parallel refine time: 0.413s speedup: 3.88x实际项目中这个加速比会被几个因素稀释多草稿 batch 的单次 forward 比单草稿 forward 更慢。GPU 显存上限决定了可并行草稿数量不是无限的。如果草稿数量太少并行化的收益不明显草稿数量太多又会挤压目标模型验证阶段的内存。所以 xPress 的工程实现本质上是在“草稿并行度”和“显存占用”之间找平衡。5. 常见问题与排查思路5.1 草稿接受率过低问题现象常见原因解决思路生成结果质量差大量 token 被拒绝草稿模型与目标模型分布差异大增加草稿模型训练或微调降低采样温度接受长度长时间为 0扩散草稿生成器没有收敛增加去噪步数或对草稿 logits 做校准接受率是 Speculative Decoding 最重要的指标之一。如果接受率太低草稿阶段白白浪费算力。建议在实验中单独记录每一轮的接受长度定位是某个 token 位置频繁失败还是整体分布偏差。5.2 并行化后显存暴涨问题现象常见原因解决思路显存不足 OOM草稿 batch 目标模型验证 batch 同时过大减少并行草稿数量开启梯度检查点推理时为显存优化模式使用 bf16速度没有提升甚至变慢显存带宽成为瓶颈用 nvidia-smi 观察显存占用和利用率适当降低草稿并行数并行细化本质上是“拿显存换延迟”。如果你的业务场景显存已经很紧张建议先用 2 个草稿做对照确认收益为正再逐步扩大并行度。5.3 推理速度不升反降问题现象常见原因解决思路token/s 反而低于普通自回归草稿每步去噪耗时太长减少num_denoise_steps或用蒸馏过的 few-step 扩散模型多 GPU 场景下通信开销大草稿和验证模型分布在多卡频繁传输把草稿模型和目标模型放在同一张卡或同一节点内使用张量并行吞吐高但首 token 延迟也高草稿阶段发生在首个回复 token 之前对首 token 阶段直接走目标模型不启动草稿流程建议用 profiling 工具拆解每个阶段的耗时草稿耗时、验证耗时、接受率、采样耗时。很多情况下瓶颈不在算法本身而在数据搬运和调度不均衡。5.4 结果不稳定问题现象常见原因解决思路同一 prompt 多次生成结果不同扩散模型采样随机性 投机解码本身带随机性设置好随机种子区分“精确复现”和“生成多样性”测试固定种子下仍不稳定CUDA 非确定性算法设置torch.backends.cudnn.deterministicTrueSpeculative Decoding 的采样过程本身具有随机性即使使用相同的 prompt多次运行结果也可能不同。如果你想复现论文里的实验结果需要固定所有随机种子并在相同硬件和库版本下运行。6. 最佳实践与工程建议6.1 草稿长度与批大小权衡草稿长度draft_len和并行草稿数量是 xPress 中最重要的两个超参数。草稿太长目标模型单次验证的 token 数量变大显存和计算开销增加。草稿太短扩散模型初始化噪声向量和多次去噪的开销占主导收益被摊薄。并行草稿数量决定多草稿 batch 的实际 batch size通常从 2 开始试验。建议做一组小规模网格搜索横轴是draft_len纵轴是并行草稿数指标是端到端 tokens/s 和接受率。6.2 KV Cache 与计算图优化目标模型验证阶段需要用到 KV Cache否则每次 forward 都从 prompt 开始算浪费严重。推荐使用transformers的past_key_values机制或者在自定义模型里维护 cache。对于扩散草稿模型去噪过程的中间变量占用较大可以在一次草稿生成完毕后立即释放避免常驻显存。使用torch.no_grad()包裹推理代码减少中间变量保存。6.3 性能指标与日志工程落地时不要只看“生成了一大段文字很快”这种直觉要量化记录draft_model_forward_time target_model_forward_time acceptance_len accepted_tokens_per_step tokens_per_second gpu_memory_usage推荐每个 decoding step 都输出一条结构化日志方便后续分析瓶颈。比如下面的简化格式import logging logging.basicConfig(levellogging.INFO) def log_step(step, draft_time, verify_time, accept_len): logging.info( step%d draft_time%.3f verify_time%.3f accept_len%d, step, draft_time, verify_time, accept_len, )6.4 安全与模型合规使用开源模型复现时注意目标模型和草稿模型各自的 License尤其是商用场景。数据准备阶段避免把敏感数据直接用于在线推理测试。所有实验应限制在授权的开发环境和测试集群中不要在生产环境直接做大规模替换。另外扩散模型生成离散 token 时建议对输出 token 做一次合法性过滤避免保留 tokenizer 中的特殊 token 或未定义 token 造成下游任务异常。7. 进一步学习路径如果想从 demo 走向真正的复现建议按这个顺序推进先用transformers自带的assisted_generation跑通一个小模型的投机解码感受草稿模型和目标模型配合的整体流程。阅读 xPress 论文原稿和官方代码关注其扩散草稿模型的网络结构、去噪步数设置、多草稿选择策略。把示例 1 中的简化投机解码循环换成真实模型加入 KV Cache。把示例 2 的扩散草稿模块替换为官方实现的 diffusion drafter。最后再做并行草稿 batch 化改造并测量不同超参下的速度收益。这个方向最考验人的地方在于论文看起来是一个很优雅的方法但工程落地时显存、batch、通信、调度都会影响最终收益。建议先用小模型跑通全流程记录各项指标后再迁移到大模型。以上内容是基于 xPress 的研究方向整理的经验笔记。由于不同版本的扩散模型和投机解码实现差异较大文中代码主要用于示意和验证调度思路实际生产环境请以你的目标模型和推理框架为准。如果后续有机会跑真实 benchmark再补充具体数据集上的效果对比。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价