面向 Diffusion Language Models 的表示空间 MMD
我们提出一种面向 Diffusion Language Models (DLMs) 的后训练方法:在冻结的预训练 DLM 的特征空间中,最小化生成分布与参考分布之间的 Maximum Mean Discrepancy (MMD)。 为估计 MMD,我们保留各个 token 位置上的上下文特征,从而在单次提取器前向中为每条序列获得多个观测。优化时,离散模型使用 policy gradients,连续模型则通过生成 latent 进行直接微分。两种情况下,直接从这些特征计算损失,都无需完整采样轨迹或联合训练辅助模型,从而实现高效后训练。 实验表明,在 OpenWebText 上,该方法在相近熵下取得更低的生成困惑度;在 GSM8K 上,取得更好的准确率-计算量权衡。在采用混合 masked-uniform diffusion 的 16B DMax-LLaDA2.0 模型上,我们在数学与代码基准上以相近或更高准确率提升了解码并行度。
论文精读
TL;DR 提出在冻结扩散语言模型表征空间最小化 MMD 的后训练方法,用单次特征提取获得多样本估计,无需完整采样或辅助模型,在生成困惑度与数学推理上取得更优质量-计算权衡。
问题
问题背景
扩散语言模型(DLMs)作为自回归生成的可扩展替代方案,支持并行解码和可控生成,但如何高效后训练以提升生成质量与推理效率仍是开放问题。
现有方法局限
传统 DLM 后训练主要依赖两类路径:
- 强化学习(REINFORCE):需要对离散 token 序列进行完整采样,梯度估计方差高,且需大量采样步数,训练开销大。
- 辅助判别器(GAN/蒸馏):需要联合训练额外模型或教师模型,训练不稳定,且可能引入模式坍塌或分布偏移。
此外,直接在 token 空间优化(如交叉熵或序列级奖励)往往难以保持生成多样性,而现有特征匹配方法多针对图像连续特征,语言模型的高维离散 token 和长序列结构使分布匹配难以有效估计。
为什么这个问题难/重要
语言生成的离散性导致梯度不可导,连续松弛或策略梯度各有缺陷;分布匹配需要在特征空间中找到适合的表示,且要从单次前向传播中高效估计 MMD,避免额外采样或辅助模型。该问题直接关系到 DLM 能否在实际部署中替代自回归模型,业界对少步解码、推理并行度、准确率-计算量权衡 高度关注。
行业类比
类似在图像生成中用 perceptual loss 替代像素级损失,本文方法在特征空间对齐分布,可类比为为文本生成提供一种轻量级“表示级奖励”,无需完整 RL 循环或强教师模型。
核心洞察
- - 在冻结的预训练 DLM 特征空间计算 MMD,实现了无需辅助模型和完整采样轨迹的分布匹配后训练。相较于传统 GAN 式方法需要联合训练判别器,以及序列级 MMD 需要冗长的采样过程,该方法直接提取上下文 token 特征,从单次前向即可获得多个观测,既降低了训练开销,又避免了辅助模型引入的训练不稳定问题,为扩散语言模型的后训练提供了更轻量且稳定的分布对齐信号。
- - token 级特征的多观测估计使 MMD 分布匹配从序列级细化到每 token 位置,单次前向传播即可获得低方差估计,显著提升样本效率。这一设计不同于以往仅用整句表示或需要多次采样来估计分布差异的方法,它为离散模型(policy gradient)和连续模型(直接微分)都提供了稳定且计算友好的训练信号,并且能扩展到 16B 参数的混合扩散模型,在提高推理并行性的同时保持或提升精度,体现了工程实用性。
方法
方法详解
输入:预训练的扩散语言模型(Discrete 或 Continuous DLM)与参考文本分布(如 OpenWebText 或 GSM8K 数据)。
关键模块:
- 表示空间选择:使用冻结的预训练 DLM 作为特征提取器,保留每个 token 位置的上下文特征向量,从而从单次前向传播中获得多个观测样本,无需额外编码器。
- MMD 估计:在特征空间中计算生成序列与参考序列分布的 Maximum Mean Discrepancy(MMD),利用核均值嵌入度量分布差异,每个 token 特征作为一个独立观测,提高估计效率。
- 离散 DLM 优化:由于离散 token 不可微,采用策略梯度(REINFORCE)直接优化采样序列的 MMD 损失,无需完整采样轨迹或联合训练的辅助模型。
- 连续 DLM 优化:连续 latent 可微,因此通过生成 latent 的直接微分优化 MMD 损失,支持 self-conditioning 的迭代细化,并可结合 bootstrapping 与迭代细化蒸馏。
输出:后训练调整后的 DLM 权重,在保持熵相近的情况下降低生成困惑度,提升条件生成的准确率与解码并行度。
方法核心是直接在冻结特征空间中最小化 MMD,将分布匹配从文本空间转移到语义特征空间,避免对抗训练的不稳定性。
与同类方法的差异点:不同于使用 KL 散度、对抗训练或需要联合训练辅助判别器的分布匹配方法,本方法利用冻结 DLM 的 token 级特征进行 MMD 匹配,无需额外模型且计算高效。
实验
实验设计
- 在 OpenWebText 上评估无条件生成,比较离散与连续 DLM;在 GSM8K 及 TinyGSM 上评估条件生成(数学推理)。
- 基线上:离散模型对比 IDLM、IDLM-REINFORCE、DiDi-Instruct;连续模型对比 ELF、ELF-PD、ELF*、ELF-GAN。
- 扩展性实验在 16B DMax-LLaDA2.0 混合 masked-uniform diffusion 上验证。
关键发现
- 在 OpenWebText 上,MMD 后训练获得更低生成困惑度,且熵与基线相当,说明没有牺牲多样性。
- 在 GSM8K 上,准确性-计算权衡 更优:相同计算预算下准确率更高。
- 在 16B 模型上,能增加解码并行度,而数学与代码 benchmark 精度持平或更高。
- MMD 损失直接基于冻结 DLM 的特征计算,无需完整采样轨迹或联合训练辅助模型,训练效率高。
对比解读
- 相比 REINFORCE 类方法,表示空间匹配避免高方差梯度估计,训练更稳定。
- 相比 ELF-GAN 等需要对抗训练的辅助判别器,本方法无需训练额外网络,避免模式坍塌风险。
- 连续模型通过直接微分生成隐变量,不需迭代细化蒸馏,简化流程。
行业影响
落地场景
该方法可直接应用于扩散语言模型(DLM)的后训练阶段,提升文本生成质量与推理能力。具体场景包括:
- 对话系统与智能助手:通过降低生成困惑度,提升回复流畅度与事实一致性。
- 代码生成与补全工具:在保持并行解码优势的同时,提高代码正确率,适用于 IDE 插件或代码托管平台的 AI 功能。
- 教育领域数学解题:利用在 GSM8K 上的准确率-计算权衡改进,可用于自动解题与讲解系统。
商业价值
核心价值在于降本增效:
- 推理成本降低:由于扩散模型本身支持并行解码,本方法进一步提升并行度,减少生成延迟,降低 GPU 消耗。
- 训练效率提升:后训练无需完整采样轨迹或辅助模型,直接基于特征匹配,训练开销小,易于迭代。
- 用户体验改善:生成质量提高(低困惑度、高准确率),带来更可靠的产品输出。
与现有工作流的集成
该方法可作为插件式后训练步骤,无缝接入现有 DLM 训练流程:
- 使用冻结的预训练 DLM 作为特征提取器,无需修改架构。
- 支持离散与连续扩散模型,覆盖主流 DLM 变体。
- 可与现有强化学习(如策略梯度)或蒸馏方法结合,进一步优化。
具体落地 Use Case
- 电商内容生成:平台需要快速生成大量商品描述、广告文案。使用本方法微调扩散语言模型,可在保证生成速度的同时,提升文案的流畅度和转化导向。
- 代码托管平台 AI 助手:例如代码补全服务,要求低延迟与高正确率。对扩散代码模型应用该后训练,可显著改善生成代码的可接受率,减少用户修改成本。
局限
- **特征空间依赖性强**:方法完全依赖冻结预训练 DLM 作为特征提取器,特征空间质量直接决定 MMD 匹配效果。若提取器未对齐目标任务(例如用通用预训练模型做代码生成),可能无法捕捉判别性特征,导致分布匹配失效。论文仅简单对比了几种特征空间,未系统分析不同层次、不同预训练目标的影响,实际部署时需额外验证。
- **离散优化稳定性不足**:对离散模型采用策略梯度估计,存在高方差和收敛慢的问题。论文虽然报告了 policy-gradient group size 等超参数消融,但未深入分析训练过程中的梯度方差、大模型下的稳定性,以及与其他低方差估计器(如 Gumbel-Softmax 松弛)的对比。在大规模模型(如 16B)上,策略梯度可能面临计算和优化瓶颈。
- **实验覆盖任务有限**:主要验证了 OpenWebText 无条件生成和 GSM8K 数学推理,虽然提到了 16B 模型上的代码基准,但缺少对话、摘要、多语言等更广泛生成任务的评估。与最新 diffusion language model 后训练方法(如 ELF 系列、DiDi-Instruct)的对比不够全面,尤其缺少在指令跟随、人类偏好对齐等实用场景下的表现,限制了结论的泛化性。