DSpark + SGLang:半自回归投机解码与置信度调度实战拆解

发布时间:2026/9/7 10:16:44
DSpark + SGLang:半自回归投机解码与置信度调度实战拆解 做 LLM 推理加速的人最近应该都在关注 DSpark 这个词。它做了一件看起来有点绕的事在不牺牲生成质量的前提下把传统自回归解码里的“串行等待”改成“按块预测 按置信度中途纠正”也就是半自回归投机解码。更关键的是这套逻辑跟前端非常流行的 SGLang 深度绑定落地时可以直接借助它的 RadixAttention、结构化前端和灵活的后处理 hook。这篇文章我从头拆一遍 DSpark 的核心思路、置信度调度的数学直觉以及它在 SGLang 里到底怎么接进去适合已经跑过 vLLM 或 SGLang 推理、想进一步压榨解码延迟的工程师阅读。我会尽量把每一步为什么这么做讲清楚代码也给了可执行的示意版本你可以直接抄去改。1. 从投机解码说起DSpark 到底在优化什么1.1 自回归生成为什么慢先算一笔账自回归解码每一回合只能生成一个 token下一个 token 必须等前面所有 token 都生成完才能开始。这个“串行”特性不是工程没做好而是模型结构决定的第 t 个 token 的 hidden state 依赖前 t-1 个 token 的 hidden state不存在严格意义上的并行空间。真正的瓶颈在 memory bandwidth。以 7B 模型为例权重参数大约 14GB按 BF16 计算。在 A100 上显存带宽大约是 2TB/s 到 3TB/s单次 decode 光把权重从显存搬到计算单元就要 5ms 左右。但每个 decode step 实际需要的浮点运算量并不高因为序列长度短注意力部分很小。换句话说GPU 大部分时间不是在“算”而是在“等数据从显存里流过来”。如果你把 batch size 拉大多个序列可以共用权重搬运每个序列的额外开销被摊薄吞吐自然会上去。但延迟问题依旧存在无论 batch 多大单条请求生成 100 个 token就得串行跑 100 次 decode。投机解码的思路就是用额外的计算或存储去换取“一次前向多验证几个 token”从而在端到端延迟上打破自回归的线性约束。1.2 投机解码的两种主流思路投机解码并不是新东西经典方案是拿一个小模型当 drafter先生成一段草稿 token比如 4 个再用大模型一次性验证这 4 个 token。验证时这 4 个 token 都作为输入进入模型并行计算各自的 logits检查与草稿是否一致。如果一致就全部接受中间有分歧就截至第一个分歧点接受其余丢弃。这个方案的关键指标叫“接受率”acceptance rate也就是草稿 token 通过大模型验证的比例。如果接受率是 0.7平均一次验证能“凭空多生成”约 1.7 个 token理论上端到端耗时就接近原来的 1/1.7。但它有个前提草稿模型要足够快、足够小否则草稿阶段省下的时间又被草稿模型自己的前向开销吃回去。另一种思路是并行独立解码例如 Medusa它不去找额外的草稿模型而是直接在原始模型上挂几个解码头一次前向同时预测当前 token 之后的多条候选路径。这种方案规避了加载第二套模型的开销但如何选择候选路径、如何处理多路径之间的依赖就成了新的工程问题。DSpark 本质上属于这一支但它改进了“多路径如何调度”的问题引入了置信度这个控制信号。1.3 DSpark 的半自回归视角把“验证”变成“调度”DSpark 的英文全称我记不太清了但它的核心主张很明确不要用固定长度的草稿块也不要在每个候选路径上浪费算力而是让草稿阶段自己判断“这一块我可以预测多远”。传统自回归是每次预测 1 个 token全并行方案是一次预测 k 个 token而半自回归介于两者之间草稿头draft head一次输出一小段 token但这段 token 不强制全接受而是通过置信度动态截断。置信度高就继续向前置信度低就回退到单步验证。换句话说DSpark 把解码过程从“生成-验证”的二元循环变成了一个有置信度反馈的闭环控制系统。这个视角的转变很重要。它意味着你不再需要精心挑选一棵猜测树也不用担心草稿模型和大模型的分布偏差太大。只要置信度标定得准调度器就能自动找到“多猜几个 token”和“少犯错误”之间的平衡点。这就是 DSpark 对投机解码最大的增量贡献。2. 置信度调度的核心机制2.1 为什么置信度本身也是个模型先问一个问题top-1 概率高代表这个 token 一定对吗在投机解码里草稿头给出的候选 token 是否被大模型接受本质上是“草稿分布”和“目标分布”是否一致的问题。如果草稿头的分布接近目标分布那么草稿给出的 top-1 token 大概率就是大模型的 top-1 token接受率就高。反过来如果草稿头没训练好或者遇到了分布外的输入即使草稿头自己给出的置信度很高大模型也可能拒绝。置信度调度做的事情就是用草稿头输出的概率值去估计“这个候选 token 被接受的可能性”。如果概率值能被正确校准那我们就可以画一条线置信度高于某个阈值的 token 被批量接受低于阈值就回退到保守的单步解码。这条线就是“精度”和“召回率”的取舍点。阈值高意味着只有很确信的 token 才会被批量通过接受率未必高因为很多置信度中等的 token 也被乖乖走了单步阈值低则会有更多错误 token 混进来虽然表面上看草稿通过率很高但大模型验证后发现分歧前面已经接受的 token 也要回滚反而浪费时间。所以阈值不能一拍脑袋定需要对草稿头做校准实验。2.2 置信度的量化方式与校准最常见的置信度是 softmax 后的 top-1 概率也就是max(softmax(logits))。它直观、计算成本低但有个毛病模型训练时通常用交叉熵使得概率分布“过度自信”所以直接拿它当接受概率会偏高。更稳的方式有两种一种是对 logits 做温度缩放后再取 top-1 概率让分布稍微平滑一点另一种是计算 top-1 和 top-2 之间的概率差。概率差大说明候选之间区分度高模型更“坚定”这类 token 一般更容易被目标模型接受。我在实验里更喜欢同时保留这两个信号top-1 概率用于判断“基础把握”概率差用于判断“候选歧义程度”。校准的方法也简单。准备一组验证 prompt记录每个草稿 token 的 top-1 概率和真实接受结果然后按概率分成若干个桶。比如概率在 0.5-0.6 的桶里统计实际接受率是多少0.6-0.7 的桶里实际接受率又是多少。如果实际接受率和概率值严重脱节说明草稿头需要重新训练或者温度参数需要调整。2.3 半自回归中的截断策略block 与置信度的联合决策有了置信度之后调度器怎么用它来决定“往前走多远”我常用的策略是“前缀截断”草稿头一次预测多个 token从第一个 token 开始逐位判断。只要当前 token 的置信度低于阈值就立刻停止扩展把前面已经高于阈值的 token 打包送去验证。这个策略等价于把块长从固定值变成了模型自适应的值因此叫半自回归其实非常贴切——它有时一口气走 3 步有时只走 1 步但不会因为单步置信度低就完全放弃批处理。也可以做“块级判断”当整块的联合置信度超过阈值时直接提交整块验证否则在块内继续寻找最长可接受前缀。块级判断的优势是减少调度器的调用次数缺点是对联合概率的估计容易偏。我的建议是块大小不超过 4 时直接用逐位截断简单且调试方便块大小超过 8 时再考虑联合置信度否则收益不明显。这里有一个容易被忽略的细节验证阶段不仅是检查“草稿 token 是否等于大模型 argmax token”还可以用随机采样方式做更柔和的验证。例如大模型给出的分布 p 和草稿分布 q 都是已知的可以按min(1, p(x)/q(x))的概率接受草稿 token。这个公式最早来自投机解码论文后来很多系统实现都用它。它允许一定的随机性能缓解草稿分布偏尖锐的问题。3. SGLang 实现思路与关键代码拆解3.1 为什么不直接选 vLLM而是 SGLang如果你搜过相关热词会发现 sglang 和 vllm 的对比很常见。两者都是高性能推理引擎但 SGLang 有几个特性对 DSpark 这种调度型解码非常友好。第一是 RadixAttention。它把 KV Cache 按前缀树复用同一 prompt 的前缀无论在预填充阶段还是 decode 阶段都能命中缓存。在投机解码里草稿块的生成和验证会反复接触相近前缀RadixAttention 能把这部分开销降到接近零。第二是 SGLang 的前端抽象。它的function装饰器允许你在 Python 层描述生成流程方便注入自定义采样逻辑。vLLM 的 sampling 参数虽然也开放但想要中途读取置信度并动态调整策略必须改 engine 内部代码维护成本高。第三是 SGLang 的底层调度器对连续 batchcontinuous batching的支持更激进。它会把不同请求的 decode step 合并成一个大 batch正好适合投机解码场景下“部分请求在验证草稿块部分请求在生成草稿”的混合状态。下表是我在同样 1 张 A100 上跑 7B 模型时记录的粗略对比指标SGLang默认vLLM默认长 prompt 前缀复用好RadixAttention一般依赖 block manager自定义 logits 后处理友好可在前端接入相对受限speculative decoding 生态正在快速补齐支持 EAGLE 等但定制成本高动态采样策略可通过 Python hook 修改需要侵入 C 层当然这不代表 vLLM 不行它在连续 batching 和调度稳定性上非常成熟。只是如果目标是快速迭代 DSpark 这类新调度算法SGLang 的灵活度更高。3.2 半自回归草稿头一次输出多个候选 token先写一个简化版的草稿头。它接收模型倒数第二层或最后一层的 hidden state然后通过一个小型 MLP 预测后续多个 token 的 logits。这与 Medusa 的做法类似但 DSpark 给它配上了置信度调度。import torch from torch import nn class DSparkDraftHead(nn.Module): def __init__(self, hidden_size: int, num_extra_tokens: int 3, vocab_size: int 32000): super().__init__() self.num_extra_tokens num_extra_tokens self.proj nn.Sequential( nn.Linear(hidden_size, hidden_size), nn.SiLU(), nn.Linear(hidden_size, num_extra_tokens * vocab_size), ) def forward(self, last_hidden: torch.Tensor) - torch.Tensor: # last_hidden: (batch, hidden_size) logits self.proj(last_hidden) batch logits.shape[0] # 返回形状 (batch, num_extra_tokens, vocab_size) return logits.view(batch, self.num_extra_tokens, -1)训练这个草稿头时输入是模型已经生成到的最后一个位置标签是接下来连续的 k 个 token。损失函数用交叉熵就行。需要注意的是草稿头不要在冻结模型的情况下一次性拿全部训练数据硬怼而是应该用“在线蒸馏”的方式让目标模型真实地生成一批 token再用这些真实 token 训练草稿头。否则草稿头会偏向训练集分布实际推理时分布一偏移置信度全部失真。3.3 置信度调度器从 logits 到 accept/reject 的循环在 SGLang 里接入 DSpark我不会建议你去改它的 C kernel而是把调度逻辑放在模型前向之后、采样之前。下面是核心循环的伪代码你可以在gen之前把 logits processor 挂上。def dspark_decode(model, prompt_ids, max_new_tokens64, tau0.6, block_size3): generated list(prompt_ids) # 简单的置信度自适应根据历史接受率调整阈值 history_accept_ratio 1.0 while len(generated) max_new_tokens: # 1. 草稿阶段用最后 hidden 预测 block_size 个候选 token last_hidden model.get_last_hidden(generated) draft_logits draft_head(last_hidden) # (1, block_size, vocab_size) draft_probs torch.softmax(draft_logits, dim-1) draft_tokens draft_logits.argmax(dim-1).squeeze(0) # (block_size,) confs draft_probs.max(dim-1).values.squeeze(0) # (block_size,) # 2. 置信度截断找最长可接受前缀 accept_prefix_len 0 for i, conf in enumerate(confs): if conf tau: accept_prefix_len i 1 else: break if accept_prefix_len 0: # 回退到普通的单步解码 new_token model.step(generated) generated.append(new_token) continue # 3. 验证阶段大模型并行验证这 accept_prefix_len 个 token candidate_ids generated draft_tokens[:accept_prefix_len].tolist() verify_probs model.get_probs(candidate_ids) accepted 0 for j in range(accept_prefix_len): # 如果草稿 token 概率足够高就接受否则按概率采样接受 target_p verify_probs[-accept_prefix_len j, draft_tokens[j]] draft_p draft_probs[0, j, draft_tokens[j]] if torch.rand(1).item() min(1.0, target_p / (draft_p 1e-6)): accepted 1 else: break generated.extend(draft_tokens[:accepted].tolist()) # 4. 根据最近 epoch 接受率调整 tau history_accept_ratio 0.9 * history_accept_ratio 0.1 * (accepted / block_size) if history_accept_ratio 0.3: tau max(0.3, tau - 0.05) else: tau min(0.9, tau 0.01)这段代码里有一个很重要的细节验证公式min(1.0, target_p / draft_p)。它是投机解码的标准接受规则。如果草稿分布和目标分布完全一致那么无论草稿 token 是什么接受概率都是 1如果草稿头对一个 token 给出了很高的概率但目标模型认为它概率极低那么 target_p 会比 draft_p 小很多token 几乎一定拒绝。还要注意model.get_probs(candidate_ids)这一步不能真的重新跑一遍完整前向。在 SGLang 里你应该复用 prefill 阶段已经算出的 KV Cache只对新增的候选 token 做增量前向。这也是 RadixAttention 在投机解码里意义重大的原因。3.4 在 SGLang 里与 reranker、dflash 等特性协同很多人在搜 SGLang 时会看到 qwen3-reranker、dflash 这些词。先说 reranker在 RAG 场景中靠向量检索拿回来的候选文档通常不够准再用 qwen3-reranker 这类重排序模型做一次精排。DSpark 和 reranker 并不冲突反而可以配合DSpark 草稿头能生成大量候选延续reranker 只对和 query 语义相关的候选块做精排再把结果喂回置信度调度器。这样就形成了“草稿-重排-验证”的三级流水线。dflash 这类名字在我见过的 SGLang 社区版本里一般指和 flash attention 融合相关的实验性调度改进。它优化的是 attention kernel 层面的数据搬运和 DSpark 的置信度调度不在同一层。如果你开了 dflash 之类的融合算子建议先验证草稿头的前向路径是否兼容因为某些融合 kernel 只支持标准的自回归输入格式遇到多 token 同时验证这种 shape 会直接回退到慢路径。经验是DSpark 的收益主要在调度层对底层 kernel 的依赖不大。无论你开不开 dflash置信度调度本身都能跑但若想发挥最大性能底层 kernel 需要支持(batch, seq_len)可变长度的批量验证否则验证阶段的长尾延迟会抵消掉草稿阶段省下的时间。4. 实测调参常见问题与排查实录4.1 关键参数怎么给初始值我通常从四组参数开始调草稿块大小 block_size、置信度阈值 tau、温度 temperature、batch size。下面是我在一台 A100 上跑 7B 模型时的经验值。参数推荐初始值调整方向说明block_size2 或 3显存充足时可升到 4块太大草稿头训练难度增加tau0.5 到 0.6接受率低就降质量变差就升需要配合校准实验temperature0.7 到 1.0低温下 top-1 概率偏高容易过自信低温时 tau 要相应调高batch_size8 到 32越大越能掩盖验证阶段的 kernel 开销小 batch 下收益不明显一个很容易踩的坑temperature 调太低时草稿头的 top-1 概率几乎全都在 0.9 以上此时如果你还保持 tau0.5那绝大多数 token 都会进入批量验证表面上草稿通过率很高但大模型验证后会发现很多分歧频繁回滚反而更慢。正确做法是低温采样场景把 tau 抬高到 0.7 以上让调度器变得“更挑剔”。4.2 接受率低问题可能不在阈值如果发现草稿块频繁被拒绝很多人第一反应是调低 tau。但接受率低更常见的原因是草稿头没训练好。我遇到过一种情况草稿头在训练集上准确率很高但换到线上 prompt 后接受率直接掉到 30%。最后定位到是领域漂移——草稿头看到的全是通用文本线上全是代码和 JSON分布完全对不上。解决思路有两个。一是做在线微调每隔一段时间用模型实际接收到的 prompt 蒸馏出一批新数据更新草稿头。二是在调度器中增加领域信号比如对代码块、表格内容采用更低的 tau因为它们本身的结构化程度高自回归模型的不确定性分布和自然语言不同。另外logits 的数值范围也需要留意。如果草稿头用的是独立初始化的 MLP它的输出 logits 方差可能和主模型不匹配softmax 后会出现“虚假的高置信度”。我建议加载草稿头后先跑几百条样本统计 top-1 概率的中位数如果中位数超过 0.95 甚至接近 0.99就要考虑对 logits 做 LayerNorm 或温度缩放。4.3 显存上涨和吞吐抖动怎么查投机解码通常会让显存占用比普通解码高因为草稿头本身是额外参数而且验证阶段需要缓存更多中间状态。我见过最夸张的情况是显存上涨了 1.5GB后来发现是多头候选路径的 logits 全部被保留在显存里没释放。解决办法是让草稿头输出的 logits 用临时变量包裹验证完成后立即释放。吞吐抖动的根源往往是验证阶段的等待。比如某个请求的草稿块长度为 3但另一个请求只有 1调度器如果强行为所有请求统一 block_size就会让短块请求白白等待。SGLang 的 continuous batching 天然能缓解这个问题但如果你是自己写 Python 循环一定要用动态截断别用固定长度。我调试时最喜欢看的三个指标是平均接受长度平均每次草稿验证接受了几个 token、草稿前向耗时、验证前向耗时。三个指标放在一起能立刻定位瓶颈。如果草稿前向耗时远大于验证耗时说明草稿头太大考虑缩小 hidden size如果验证耗时偏高考虑是不是把候选 token 全部重新计算了前缀而不是复用 KV Cache。5. 一些个人体会和后续可以玩的方向DSpark 最打动我的地方是它把“能不能多预测几个 token”这个问题从纯模型结构层面转移到了调度层面。以前我们总觉得投机解码就是找一个好的草稿模型DSpark 却告诉你即使草稿头不完美只要置信度调度得够好依然能在正确率和速度之间拿到不错的折中。这种视角对工程实现非常友好因为你在生产环境里没法保证草稿头永远不漂移但你可以让调度器根据即时反馈自动调整行为。后续我觉得有两个方向值得尝试。一个是把置信度阈值从全局参数升级成按 token 位置、按 prompt 类型动态预测的小模型相当于用一个 meta model 去学“什么时候该冒险”。另一个是把它和 beam search 或 tree search 结合草稿头一次生成多条候选路径置信度调度负责判断哪些路径值得保留再统一交给目标模型验证这样等于把投机解码从直线扩展成了树形理论上能进一步提升每次验证的收益。如果你最近正要给 SGLang 接自定义解码策略建议先把 DSpark 的调度循环写成独立的 Python 模块不要一上来就动 C 层。等收益验证清楚再考虑下沉到 kernel 层优化也不迟。毕竟调度算法的迭代速度远比 kernel 优化快得多。