DiffusionGemma 技术报告:开源离散扩散大语言模型 DiffusionGemma Technical Report
把 Gemma 4 微调成开源离散扩散模型,单 H100 上约 1500 tokens/s。
前置知识
自回归语言模型 (Autoregressive LLM)
自回归模型把序列联合概率 $p(x)$ 分解为从左到右的条件概率连乘 $p(x)=\prod_{i} p(x_i\mid x_{<i})$,每次只预测下一个 token,GPT、Gemma 等都属此类。优点是训练目标与似然严格对应,缺点是生成必须串行解码,无法回看未来 token 做修正。
本文的全部出发点就是 AR 模型在低并发推理下被显存带宽卡住、无法并行生成;理解 AR 的串行瓶颈才能理解扩散模型为什么要重新设计生成机制。
离散扩散 / 多项式扩散 (Discrete / Multinomial Diffusion)
扩散模型把生成建模成马尔可夫链:前向过程把干净 token 腐蚀成均匀噪声(每个 token 以概率 $t$ 换成词表随机 token),反向过程训练网络 $p_\theta$ 预测噪声状态下原始 token 并去噪。多项式扩散允许 token 间互相转移,使模型能在后续步骤纠正错误 token。
DiffusionGemma 的核心就是把整段 256 个 token 的画布并行去噪,这是它实现高 TPF(每前向 token 数)的机制基础。
混合专家模型 (Mixture-of-Experts, MoE)
MoE 把前馈网络拆成多个专家,每个 token 只激活少数几个专家(本文是 8/128 + 1 个共享专家),从而用较小的激活参数量(3.85B)撑起较大的总参数量(25.2B)。代价是单请求服务时,专家权重搬运成为显存带宽瓶颈。
本文专门分析了 MoE 在扩散并行解码下的专家激活爆炸问题(画布级平均激活 84 个专家 vs 单 token 的 8 个),这是理解其推理瓶颈与优化方向的关键。
多 token 预测 / 投机解码 (MTP / Speculative Decoding)
MTP 或投机解码用一个小模型或预测头一次性草拟约 8 个候选 token,再交给大模型一次前向验证,能实现每前向 3-6 token 的吞吐。它是当前 AR 模型加速的主流手段。
论文拿它当最强的 AR 加速基线,证明扩散模型每前向 ~20 token 仍然大幅超越配备 MTP 的 AR 模型。
自条件化 (Self-Conditioning)
自条件化把模型上一步的预测(softmax 概率经嵌入矩阵投影成连续向量)作为额外输入回灌给自己,让去噪过程能记住之前的判断,类似图像扩散里预测 $x_0$ 再喂回的做法。本文用一个仅 7.8M 参数的小 FFW 块来实现。
这是让离散扩散在不增加去噪步数的情况下收敛得更好、并显著减少 token 结巴退化的关键技巧。
强化学习对齐 (RL for LLM)
RL 微调让模型在线生成、再用奖励模型或规则奖励打分、按策略梯度更新,从而提升指令遵循、数学、编码等能力,是 Gemma 4 等 AR 模型后训练的标准环节。
本文的 SD·RL 阶段把 RL 对齐和采样器蒸馏合在一个在线阶段里同时做,是其方法新颖性的来源之一。
研究动机
当前大语言模型几乎都是自回归的,在单请求或低并发推理时是 memory-bound:把模型权重和 KV cache 从显存搬到加速器的时间远超实际计算时间,算力被严重浪费。即便用最先进的投机解码或多 token 预测,草拟长度 8 时也只能换来每前向 3-6 个 token,而且 AR 草稿器本身受串行限制、并行草稿器的接受率还会随草稿位置衰减。批量请求能靠 batching 拉高吞吐,但单用户延迟依然被卡死。与此同时,现有文本扩散模型要么锁在闭源 API 后(如 Gemini Diffusion、Mercury),要么虽开源但推理能力、多模态理解较弱,或干脆没兑现极速承诺——速度、智能、开放三者始终无法兼得。
本文的目标是本文要造一个同时满足高智能、极致低延迟、完全开放权重这三个条件的文本扩散模型,并在质量-速度 Pareto 前沿上同时超越 AR 家族(含 MTP)和所有现有文本扩散模型:在单块 NVIDIA H100 上跑到约 1,500 tokens/s 的输出速度,同时保留原模型的长上下文、多模态输入与思考模式能力,并希望同一套权重还能切回纯 AR 模式,从而支持按延迟约束与任务复杂度动态路由的混合解码方案。
与已有工作不同的是,它的独特切入点是不从零预训练,而是把已经 post-train 完成的 Gemma 4 26B A4B MoE 直接微调成扩散模型,用不到原始 AR 训练 10% 的 token 预算,就把同一个 transformer 骨架同时改造成因果编码器(处理上下文并维护可增量追加的 KV cache)和双向解码器(去噪画布)。这种 AR 权重复用 + 编码器-解码器结构反转 + 两阶段 SFT/SD·RL 训练的组合,恰好填补了开放、智能、极速三者长期并存的空白。
核心方法
直觉上,文章把写一段话从逐字写改成先撒一画布噪声再反复擦改:模型一次并行处理 256 个 token,从均匀噪声出发迭代去噪,平均约 12 步(最多 48 步)收敛成一段连贯文本。技术路线上,DiffusionGemma 是一个共享权重的 encoder-decoder transformer:因果编码器把系统指令、用户输入和已生成画布编进 KV cache(从而继承多模态与长上下文能力),双向解码器在画布内做双向注意力、对 KV cache 做交叉注意力来去噪;每个画布去噪完就把它的 KV 追加进缓存,再开始下一个画布,形成块自回归(block-AR)。训练分两阶段:SFT 教双向去噪,SD·RL 同时做奖励最大化与采样器蒸馏。
核心创新有三点。第一是因果编码 + 双向解码的结构反转——不同于 BART/T5 那种双向编码器+因果解码器,它反过来用因果编码器维护可增量追加的 KV cache(避免对不断增长的上下文重编码),用双向解码器做扩散,从而把扩散与 AR 自然融合。第二是把奖励驱动的质量提升和采样器蒸馏(把高质量长轨迹压进少步)塞进同一个 SD·RL 在线阶段,一个梯度更新同时优化两件事,并通过自适应停止诱导出隐式课程学习。第三是熵有界采样 + 温度退火 + 自适应停止这套采样器,让模型按任务难度动态分配算力(代码任务步数少、自然语言多),并把 TPF 从 SFT 的约 5 拉到近 20。
方法步骤详情
推理流程为:(1) 上下文编码——因果注意力把 prompt、系统指令与历史画布编进 KV cache $H$;(2) 画布初始化——从词表均匀采样 256 个 token 作为 $x_1$,自条件化信号 $z_1=0$;(3) 去噪循环(最多 $N=48$ 步)——解码器输出 logits $L_t$,按线性退火温度 $\tau_t$(0.8→0.4)做 softmax 得 $\hat{p}_0$,再按熵从小到大排序、在累计熵预算 $b=0.1$ 内接受这些 token,其余重新均匀加噪,并更新 $z_{t-\Delta t}=\mathrm{FFW}(\hat{p}_0 E)$;(4) 自适应停止——画布平均熵 $\leq 0.005$ 且连续两步 argmax 预测相同时提前终止;(5) 画布固化——把 $\hat{x}_0$ 经编码器追加进 KV cache,开始下一画布直到结束符。训练侧 SFT 用交叉熵 $\mathcal{L}=\sum_i\log p_\theta(x_0^i\mid x_t,z_t,H)$,SD·RL 用在线教师轨迹联合优化奖励与步数压缩。
技术新颖性
新颖性体现在:(1) 首次证明用 AR post-trained 权重热启动、加上不到 10% token 预算的微调,就能得到前沿级文本扩散模型,绕开了天价扩散预训练;(2) SD·RL 把传统上需要分阶段做的对齐与蒸馏统一成一个在线目标,且熵下降触发的自适应停止天然形成课程,使延长训练在奖励平台期之后仍能持续提速;(3) 采样器层面引入熵有界接受(类似 MaskGIT)+ 温度退火 + 自适应停止的组合,把每前向 token 数从 SFT 的约 5 提到近 20;(4) 双向解码让模型具备 AR 所没有的局部自纠错和动态测试时算力能力;(5) 同一套权重可无缝切回纯 AR 模式,为按延迟与难度动态路由的混合解码铺路。
实验结果
图1 的帕累托图上,DiffusionGemma 把前沿推到质量约71、速度约1500 TPS,同时超过整个 Gemma 4 AR 家族(含 MTP,约300 TPS)和所有现有扩散模型;闭源 Mercury 2 质量略高(约77)但仅约600 TPS,慢约2.5倍。它平均约20 TPF,单 H100 FP8 上约1479 TPS(thinking),相对 AR 的204 TPS 与 AR+MTP 的303 TPS 分别提升7.1倍与4.8倍。质量上 TD 模式相对 AR 有损失(GPQA 73.2 vs 82.3、LiveCodeBench-v6 69.1 vs 77.1),但换来近5倍吞吐。SD·RL 在组合 GPQA+LCB 上提升10分、TPF 从5翻到近20、输出缩短约2倍,总前向次数小于 AR 基线的5%。每步处理256倍 token 但仅慢3.2倍,瓶颈是 MoE(4.3×)、采样(3.06ms vs 0.56ms)、注意力(4×)。下游方面数独全量微调 >85%(基线0%);批次上 <32用户时 per-user 与总吞吐都赢 AR+MTP,超过约32后被 AR 反超。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| 单请求解码速度 (H100 FP8, batch 1) | 输出 tokens/s (TPS) | 约 1,479 TPS (thinking) | Gemma 4 AR 204 TPS / AR+MTP 303 TPS | 7.1× / 4.8× |
| 每前向 token 数 (TPF) | TPF | 约 19.74 | AR 约 1.0 / MTP 约 1.4 | 约 20× / 约 14× |
| SD·RL 综合质量 (GPQA-Diamond + LiveCodeBench-v6 均值, thinking) | 平均分 | SD·RL 后检查点 | SFT 检查点 | +10 分;TPF 从 5 提到约 20 (4×) |
| GPQA-Diamond (thinking) | 准确率 % | 73.2 | Gemma 4 AR+MTP 82.3 | 质量损失约 9 分换取约 5× 速度 |
| 数独求解 (LoRA rank 8 微调) | 谜题级准确率 % | 84.4 | 未微调基座 0.0 | 从 0% 到 84.4% |
局限与改进
作者承认的局限:(1) 相对 AR 基线存在质量损失,原因是绕开了原生扩散预训练、SFT 阶段偏短、SD·RL 显式追求低延迟而牺牲渐近性能,并继承了 AR 架构与数据/优化选择;(2) 输出过于简洁,无法借助更长思维链获得质量红利;(3) 偶发 token 结巴(如反复重复同一个常见 token),超低步数下生成稳健性下降;(4) 多模态任务里偶尔漏写闭合 thought 标签,导致 MMMU-Pro 上思考分(54.3)反而低于非思考分(66.0);(5) 高并发(>约 32 用户)时因每 token 算力更高被 AR 反超。我自己的观察:质量损失在能力关键型任务(GPQA 掉约 9 分)上不可忽视;简洁性对需要长 CoT 的难题是结构性短板;采样仍用标准 PyTorch 而非手写 kernel,高 batch 下还有提速空间;多模态与多语言评测覆盖不足。
独立分析的弱点
第一,质量天花板被 AR 热启动拖累——绕开原生扩散预训练、SFT 偏短,导致能力任务上明显落后 AR 基线,改进方向是做原生扩散预训练或大幅延长 SFT。第二,简洁性是双刃剑,对依赖长思维链的数学/竞赛题不利,可通过在 SD·RL 奖励里加入长度适当项来调节。第三,高 batch 吞吐劣于 AR,MoE 专家爆炸(84 vs 8)是主因,密集架构或对采样做 top-$k$ 截断可缓解,需要专门针对 batch>32 优化。第四,多模态 thought-tag 漏写是工程 bug,应修复后重测 MMMU-Pro 以反映真实能力。第五,采样未用手写 kernel,自条件化 matmul 与全画布 softmax 在 262k 词表维度上较慢,专门的融合 kernel 可进一步提速。
未来方向
作者明确的方向包括:原生扩散预训练、高并发真实流量下的吞吐优化、混合 diffusion-AR 路由(按延迟约束与任务复杂度动态切换)、以及面向扩散的采样算法创新。基于本成果可延伸的方向:把熵有界 + 自适应停止采样器迁移到其他文本扩散模型;研究 SD·RL 的隐式课程能否推广到 AR 的 RL;用更强的测试时算力缩放(更大的 $N$)补回质量损失;探索扩散与工具调用/agent 工作流的结合(其用算力换内存的特性对长上下文 agent 友好);以及把 AR/Diffusion 双模式做成统一路由服务。
复现评估
复现友好度很高:权重以 Apache 2.0 完全开放,并在 HuggingFace Transformers 与 vLLM 提供参考推理实现;Table 2 给出全部采样器超参(画布 256、最大去噪步 $N=48$、熵预算 $b=0.1$、停止阈值 0.005、温度线性退火 0.8→0.4、平均 12 步),Appendix C 报告训练细节。配套开源了基于 Hackable Diffusion 的 SFT 工具包与 LoRA 配方,2×A100 80GB 消费级硬件即可做下游微调。算力门槛相对低(不到原 AR 训练 10% 的 token),难点主要在 SD·RL 在线教师轨迹生成与奖励设计需要一定工程投入。数独案例从 0% 到 84.4% 可独立验证,整体可复现性强。
论文图表