← 返回 2026-08-04

DiffusionGemma 技术报告:开源离散扩散大语言模型 DiffusionGemma Technical Report

DiffusionGemma Team, Adrien Ali Taïga, James Assiene, Daniele Calandriello, Rahma Chaabouni, João Gante, Tamara von Glehn, Nate Keating, Chris Knutsen, Martin Kukla, Tianlin Liu, Ivan Lobov, Ofir Nabati, João Gabriel Oliveira, Nicolas Perez-Nieves, Nastasia Prutianova, Bobak Shahriari, Jean Tarbouriech, Pavel Tyletski, Çağlar Ünlü, Cindy Wu, Glenn Cameron, Jerome Connor, Sertan Girgin, Maarten Grootendorst, Alon Levkovitch, Eliya Nachmani, Omar Sanseviero, Piotr Stanczyk, Quentin Berthet, Andrew Campbell, Clément Crepy, Valentin De Bortoli, Arnaud Doucet, Romuald Elie, Alexandre Galashov, Klaus Greff, Alexis Jacq, David Ruhe, Yu-Han Wu, Sebastian Flennerhag, Brendan O'Donoghue, George Scrivener, Shantanu Thakoor 📅 2026-07-31 👍 36 2026-08-09 18:30
开源大模型 强化学习对齐 推理加速 文本扩散 混合专家MoE 知识蒸馏 离散扩散 语言模型

把 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 模式,为按延迟与难度动态路由的混合解码铺路。

Overview of our two-stage training pipeline that converts an autoregressive model (Gemma 4 26B A4B) into a text diffusion model (DiffusionGemma).
Figure 2: Overview of our two-stage training pipeline that converts an autoregressive model (Gemma 4 26B A4B) into a text diffusion model (DiffusionGemma).
Stylized example of discrete diffusion probability paths and parallel sampling trajectory.
Figure 3: Stylized example of discrete diffusion probability paths and parallel sampling trajectory.
The DiffusionGemma generation pipeline.
Figure 4: The DiffusionGemma generation pipeline.
Adaptive stopping enables DiffusionGemma to dynamically adjust its number of denoising steps to task complexity and domain.
Figure 5: Adaptive stopping enables DiffusionGemma to dynamically adjust its number of denoising steps to task complexity and domain.
Evolution of downstream performance during SFT.
Figure 6: Evolution of downstream performance during SFT.
Downstream performance during SFT improves log-linearly with training progress.
Figure 7: Downstream performance during SFT improves log-linearly with training progress.
SD·RL training simultaneously increases average reward and reduces the effective denoising steps of the online teacher.
Figure 8: SD·RL training simultaneously increases average reward and reduces the effective denoising steps of the online teacher.

实验结果

图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 反超。

Parameter counts.
Table 1: Parameter counts.
Diffusion denoising and recommended sampler hyperparameters.
Table 2: Diffusion denoising and recommended sampler hyperparameters.
Performance comparison across benchmarks and decoding modes.
Table 3: Performance comparison across benchmarks and decoding modes.
Decoding efficiency metrics.
Table 4: Decoding efficiency metrics.
Sudoku accuracy of the finetuned model with LoRA rank 8.
Table 5: Sudoku accuracy of the finetuned model with LoRA rank 8.
Pareto plot of quality versus output decoding speed comparing DiffusionGemma to the Gemma 4 model family and other diffusion models.
Figure 1: Pareto plot of quality versus output decoding speed comparing DiffusionGemma to the Gemma 4 model family and other diffusion models.
SD·RL significantly advances the quality-speed Pareto frontier.
Figure 9: SD·RL significantly advances the quality-speed Pareto frontier.
Performance vs. number of denoising steps N, without adaptive stopping or temperature annealing.
Figure 10: Performance vs. number of denoising steps N, without adaptive stopping or temperature annealing.
Per-step GPU time breakdown.
Figure 11: Per-step GPU time breakdown.
Trade-off between total and per-user throughput of the Gemma 4 AR model (with and without MTP) and DiffusionGemma.
Figure 12: Trade-off between total and per-user throughput of the Gemma 4 AR model (with and without MTP) and DiffusionGemma.
Performance by capability area and output speed.
Figure 13: Performance by capability area and output speed.
Sudoku Downstream SFT performance.
Figure 14: Sudoku Downstream SFT performance.
Denoising trace for the math reasoning problem.
Figure 15: Denoising trace for the math reasoning problem.
查看结构化数据
任务指标本文基线提升
单请求解码速度 (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% 可独立验证,整体可复现性强。