← 返回 2026-08-13

离散扩散的单纯形松弛 Simplex Relaxation for Discrete Diffusion

Jinya Sakurai, Patrick Pynadath, Satoshi Hayakawa, Jaehong Yoon, Xulei Yang, Nancy F. Chen, Xun Xu 📅 2026-08-11 👍 6 2026-08-18 18:30
Dirichlet分布 Rao-Blackwell化 文本生成 生成模型 离散扩散

以Dirichlet辅助单纯形变量精确增广均匀离散扩散,在不改前向过程下同时升级训练目标与采样器

前置知识

离散扩散模型

在离散数据(文本token、DNA碱基等符号序列)上做生成建模的扩散模型。前向过程按腐蚀核逐步把干净样本打乱,反向过程训练一个去噪网络逐步还原。按腐蚀核分为掩码扩散(token被替换为特殊掩码态)与均匀扩散(token被替换为词表上的分布),腐蚀核的选择决定了中间状态空间与反向预测问题的形式。

Simplax 正是在均匀离散扩散的基础上做增广,不先掌握腐蚀核、前向/反向过程的基本设定,就无法理解它要增强的对象与约束。

均匀扩散与前向边际

均匀扩散把 token 以概率 $1-\alpha_t$ 替换为词表上的类别分布 $\pi$,各类别对称、无特殊掩码态。给定干净样本 $x$,噪声态边际为 $q(z_t|x)=\mathrm{Cat}(z_t; p_t)$,其中 $p_t=\alpha_t x+(1-\alpha_t)\pi$,$\alpha_t$ 从 1 衰减到 0;反向后验 $q(z_s|z_t,x)$ 有闭式解。

$p_t$ 与反向后验是本文全部推导的出发点:辅助变量 $w_t$ 的均值锚定在 $p_t$ 上,反向桥 $\rho_{s|t}$ 也由 $p_t$ 与 $\pi$ 组合而成。

Dirichlet分布

定义在概率单纯形 $\Delta_{K-1}$($K$ 维非负、和为 1 的向量集合)上的连续分布,由浓度向量控制:均值是归一化的浓度向量,总浓度越大分布越集中。它是范畴分布的共轭先验:从 $\mathrm{Dir}(\eta p)$ 取一次范畴样本即得 $\mathrm{Cat}(p)$,这一对耦合是本文增广的构件。

Simplax 的辅助变量取自平移 Dirichlet,即 $q(w_t|z_t,x)=\mathrm{Dir}(\eta_t p_t + z_t)$;理解浓度参数 $\eta$ 才能明白它如何控制松弛态的散布与实验中的超参选择。

Rao-Blackwell化

统计中的降方差技巧:把对某随机变量的蒙特卡洛平均换成其解析期望。本文中,对辅助解码 $\tilde{z}_t\sim\mathrm{Cat}(w_t)$ 求平均的离散反向 KL 可精确算出闭式(两个内积之和),等价于 Rao-Blackwell 化后的目标,从而绕开对 Dirichlet 混合分布直接求 KL 的不可解困难。

这是全文的技术枢纽:没有这一步,单纯形桥上的 KL 无解析式,方法只能停留在不可解的原始形式(式 13)。

NFE 与生成困惑度-熵前沿

NFE(函数评估次数)指采样时去噪网络的前向调用次数,衡量推理预算,少步采样是离散扩散的重要卖点。生成质量常用 Gen PPL-ENT 前沿评估:Gen ENT 是生成样本的一元熵(多样性),Gen PPL 是 GPT-2 Large/XL、Llama-2 7B 等外部模型给出的困惑度(流畅度),扫温度得到权衡曲线,同等熵下 PPL 越低越好。

论文的文本实验全部按 NFE 分档报告,并以数据熵 5.44 nats 为参照线,读懂这对指标才能理解 Figure 3/4 与 Table 1 的对比逻辑。

研究动机

离散扩散用腐蚀核定义生成过程:掩码扩散(MDLM 等)把 token 吸收到特殊掩码态,训练采样都直观;均匀扩散则把 token 对称地替换为词表分布 $\pi$,不引入特殊状态,近来被证明在引导、token 重复修订、少步生成、自我纠错与 scaling 上有独特价值(Schiff et al. 2025、Sahoo et al. 2025 等)。但均匀扩散的训练目标与反向转移完全通过范畴状态表达:反向更新直接在采样的离散 token 之间进行,训练损失也只以离散态为唯一锚点,中间没有额外的概率结构可资利用。已有的增强路线各有局限:要么替换主生成态——FLM 在 one-hot 状态上做欧氏去噪、DDSM/Dirichlet Flow Matching 把单纯形当主状态、CADD/CANDI 做离散-连续混合;要么只针对掩码扩散——如 VADD 在掩码去噪中引入高斯隐变量、Di4C 用乘积模型混合、CoDD 耦合概率电路。没有一个方法在保持均匀范畴腐蚀过程完全不变的前提下,系统性地丰富其训练目标与采样器。

本文的目标是本文目标可以概括为:在不改变均匀离散扩散前向腐蚀过程 $q(z_t|x)=\mathrm{Cat}(z_t;\alpha_t x+(1-\alpha_t)\pi)$ 的前提下,为其训练目标与反向采样引入更丰富的概率结构,并保证推导严格可解。具体拆成三件事:其一,构造一个以原范畴过程为边际的精确增广层级;其二,从该层级导出可解析计算(Rao-Blackwell 化闭式)的反向桥训练目标,并给出连续时间极限;其三,从同一层级导出随机祖先采样器。最终在两组实验上验证:OpenWebText 无条件生成应改进 Gen PPL-ENT 权衡(NFE=16 到 1024 多档预算);Sudoku 上只训 30 提示数独,跨提示数(40 至 17,含唯一解下限的 17 提示)求解准确率与无条件生成有效率应全面超过 MDLM、UDLM、Duo、FLM、CANDI、LangFlow、S-FLM 等对比方法。

与已有工作不同的是,独特切入在于辅助变量的地位:Simplax 引入单纯形变量 $w_t$,但它不是新的主生成态,而是精确挂在范畴态上的辅助桥。条件分布取平移 Dirichlet $q(w_t|z_t,x)=\mathrm{Dir}(\eta_t p_t+z_t)$,可证明其边际恰为 $q(w_t|x)=\mathrm{Dir}(\eta_t p_t)$,且解码器精确为 $q(z_t|w_t)=\mathrm{Cat}(z_t;w_t)$——增广层级把原均匀扩散作为范畴边际完整保留。这与把单纯形当主态的 DDSM/Dirichlet FM、在掩码扩散中加高斯隐变量的 VADD、以欧氏几何替代离散态的 FLM 都不同。为绕开单纯形反向桥(Dirichlet 混合)KL 不可解的困难,作者对辅助范畴解码 $\tilde{z}_t\sim\mathrm{Cat}(w_t)$ 求期望得到解析闭式;同时坚持让独立采样的 $z_t$ 作为去噪网络输入,保住 token embedding 查表的工程优势,避免每位置一次词表级稠密矩阵乘 $w_t^\top E$。

核心方法

直觉上,均匀扩散的噪声态 $z_t$ 是单纯形顶点上的 one-hot 向量,反向桥只由这一个顶点锚定,信息单一。Simplax 给每个 $z_t$ 配一个单纯形内点 $w_t$:训练时按前向边际采 $w_t\sim\mathrm{Dir}(\eta_t p_t)$,再采 $z_t\sim\mathrm{Cat}(w_t)$;$w_t$ 的均值是扩散时刻边际 $p_t=\alpha_t x+(1-\alpha_t)\pi$,加上的 one-hot $z_t$ 把松弛态锚回离散 token(浓度 $\eta_t$ 主实验取常数 0.01)。网络仍只看 $z_t$,预测 $\hat{x}_\theta=f_\theta(z_t,t)$。训练损失是离散反向 KL 对辅助解码 $\tilde{z}_t\sim\mathrm{Cat}(w_t)$ 的期望,有 Rao-Blackwell 闭式 $\bar{\mathcal{L}}=\langle w_t,\log\hat{p}_t-\log p_t\rangle+\langle\rho_{s|t}(x,w_t),\log p_s-\log\hat{p}_s\rangle$,其中 $\rho_{s|t}$ 是给定 $w_t$ 的范畴型反向后验。采样时维护 $(z_t,w_t)$ 二元组:从 $w_{t_N}\sim\mathrm{Dir}(\eta\pi)$、$z_{t_N}\sim\mathrm{Cat}(w_{t_N})$ 出发,每步由 $\rho_{s|t}(\hat{x}_\theta,w_t)$ 采 $z_s$ 作下一步输入,再由 $\mathrm{Dir}(\eta_s\hat{p}_s+z_s)$ 采 $w_s$,到 $t=0$ 输出 $z_0$。全程前向范畴过程一字未改。由于增广把原扩散保留为范畴边际,这套改动只作用于训练目标与反向转移的参数化,数据分布与腐蚀核都不变;$w_s$ 每步重新围绕 $\hat{p}_s$ 集中,本质上是用一个连续的桥变量给离散反向过程配上更细粒度的匹配信号,而采样本身仍保持随机祖先形式。

核心创新是精确的 Dirichlet-范畴增广。命题 1 证明该层级同时满足:边际 $q(w_t|x)=\mathrm{Dir}(\eta_t p_t)$、精确解码器 $q(z_t|w_t)=\mathrm{Cat}(z_t;w_t)$、范畴型反向后验 $q(z_s|w_t,x)=\mathrm{Cat}(\rho_{s|t}(x,w_t))$,以及 Dirichlet 混合形式的单纯形桥。与已有方法的本质区别在于变量的角色:Di4C 用乘积模型混合、VADD 在掩码去噪中引入高斯隐变量、CoDD 耦合概率电路、CANDI/CADD 做离散-连续扩散、FLM 在 one-hot 态上做欧氏去噪——这些方法或替换或改变主生成态;而 Simplax 的单纯形变量只是精确辅助桥,原范畴过程作为边际原封不动。另一要点是把不可解的 Dirichlet 混合 KL(式 13)转化为对辅助范畴解码 $\tilde{z}_t\sim\mathrm{Cat}(w_t)$ 取期望再解析边际化,即 Rao-Blackwell 化,得到命题 2 的闭式;由于 $z_t$ 与 $\tilde{z}_t$ 给定 $w_t$ 条件独立,网络输入 $z_t$ 不被边际化,仍保留离散输入。

方法步骤详情

第一步,构造层级:按式(6)因子化 $q(x)q(z_s|x)q(w_s|z_s,x)q(z_t|z_s)q(w_t|z_t,x)$,令 $q(w_t|z_t,x)=\mathrm{Dir}(w_t;\eta_t p_t+z_t)$,$\eta_t>0$ 为浓度参数。第二步,采样增广态:每步取时间对 $(s,t)$,由前向边际采 $w_t\sim q(w_t|x)$,再采 $z_t\sim\mathrm{Cat}(w_t)$。第三步,预测与损失:算 $\hat{x}_\theta=f_\theta(z_t,t)$,按闭式 $\bar{\mathcal{L}}=\langle w_t,\log\hat{p}_t-\log p_t\rangle+\langle\rho_{s|t}(x,w_t),\log p_s-\log\hat{p}_s\rangle$ 反传,其中 $\rho_{s|t}=p_s\odot(\alpha_{t|s}\,w_t\oslash p_t+(1-\alpha_{t|s})\langle w_t,\pi\oslash p_t\rangle\mathbf{1})$;命题 3 给出连续时间极限 $\ell_{ct}$,含 $\lambda(t)=-\frac{d}{dt}\log\alpha(t)$。第四步,随机祖先采样:初始化 $w_{t_N}\sim\mathrm{Dir}(\eta_{t_N}\pi)$、$z_{t_N}\sim\mathrm{Cat}(w_{t_N})$;每步算 $\hat{x}_\theta=f_\theta(z_t,t)$,采 $z_s\sim\mathrm{Cat}(\rho_{s|t}(\hat{x}_\theta,w_t))$ 作下一步输入,再采 $w_s\sim\mathrm{Dir}(\eta_s\hat{p}_s+z_s)$,从 $t_N=1$ 迭代到 $t_0=0$ 输出 $z_0$。端到端只需范畴与 Dirichlet 采样,无需新数值求解器。整个流程的输入输出很规整:训练侧输入 $(w_t,z_t,t)$,输出对 $\hat{x}_\theta$ 的梯度;采样侧输入起点分布 Dir(ηπ),输出最终范畴样本 $z_0$。超参方面主实验取常数 $\eta_t\equiv0.01$;$z_s$ 在下一步仍按整数 token 索引做嵌入查表,$w_s$ 只驻留在 CPU/GPU 上的采样器状态里,不进网络。自条件等增强项按 Figure 3 的诊断结论配置。

技术新颖性

技术新颖性有三层。第一,几何角色上把单纯形当桥而非舞台:过去 Dirichlet 类方法(DDSM、Dirichlet Flow Matching、基于 categorical SDE 与 Cox-Ingersoll-Ross 动力学的单纯形扩散)把单纯形态作为主生成对象,改变生成过程本身;Simplax 的 $w_t$ 由 shifted Dirichlet 精确耦合到 $z_t$,生成过程的边际仍是原扩散,可无痛嫁接到已有均匀扩散框架。第二,目标推导上用辅助解码平均加 Rao-Blackwell 化绕开死路:直接匹配单纯形桥需求 Dirichlet 混合的 KL,无解析式;对 $\tilde{z}_t\sim\mathrm{Cat}(w_t)$ 求期望后闭式恰为两个内积,且命题 3 证明其连续时间极限与 UDLM 目标构成单纯形松弛类比,理论谱系清晰。第三,工程上坚持 $z_t$ 作网络输入:若以 $w_t$ 为输入,每位置需做 $w_t^\top E$ 的词表级稠密矩阵乘,嵌入查表优势尽失;Figure 3 消融显示同等 Gen ENT 下 $z_t$ 输入的 Gen PPL 更低,方法与工程决策互相印证。

Reverse-bridge matching, illustrated for K = 3: one-hot states are vertices of ΔK−1 and wt is an interior point. (a) The corrupted state zt is both the denoiser input and the sole bridge anchor, so a single bridge is matched per sample. (b) Simplax samples (wt, zt) jointly; zt remains the denoiser input, while extra draws ezt ∼ Cat(wt) anchor all K bridges, which are matched at once and marginalized exactly to give (15).
Figure 1: Reverse-bridge matching, illustrated for K = 3: one-hot states are vertices of ΔK−1 and wt is an interior point. (a) The corrupted state zt is both the denoiser input and the sole bridge anchor, so a single bridge is matched per sample. (b) Simplax samples (wt, zt) jointly; zt remains the denoiser input, while extra draws ezt ∼ Cat(wt) anchor all K bridges, which are matched at once and marginalized exactly to give (15).
Graphical model corresponding to the factorization in (6).
Figure 2: Graphical model corresponding to the factorization in (6).

实验结果

文本实验:在 OpenWebText(GPT-2 BPE、$|V|=50{,}257$、序列长 1024、179M 参数 DiT、Adam lr $3\times10^{-4}$、batch 512、1M 迭代)上扫 15 档温度,选生成熵最接近数据熵 5.44 nats 的工作点。Table 1:NFE=16 时 Simplax 的 GPT-2 Large PPL 为 90.5,低于最强基线 CANDI 的 97.2、MDLM 的 117.9 与 UDLM 的 186.0;NFE=1024 时降到 45.1,比 MDLM 的 55.1 低约 18%;NFE=128 时 GPT-2 L/XL 下最优(56.9/58.9),Llama-2 7B 下以 31.4 略逊于 LangFlow 的 30.0。设计诊断(Figure 3,50k 步短训):自条件仅在 NFE=128 有益;$z_t$ 输入优于 $w_t$ 输入;UDLM 800k 到 Simplax 200k 的初始化在匹配 1M 预算下改进前沿。数独实验:所有模型只训 30 提示数独(48,000 训练题、180 token 序列、骨干 8 层/隐藏 512、25.21M-28.59M 参数、20k 步),跨提示数评估(Table 2):40 提示 98.55%、35 提示 91.05%、30 提示 61.75%、25 提示 25.90%、20 提示 8.80%、17 提示 1.20%,全部第一;30 提示下最强基线 Duo 为 48.85%,领先 12.9 个百分点;无条件生成有效率 95.85%,比 Duo 的 80.95% 高约 15 个百分点,自回归基线仅 8.15%。跨评估器看,NFE=16 时 Llama-2 7B 指标 Simplax 为 49.3,对比 CANDI 的 56.0、MDLM 的 63.8;表中加粗/下划线标注最优/次优,Simplax 在 9 个文本指标列中拿下 8 个最优。数独侧的跨密度迁移尤其值得注意:所有模型都没见过 30 提示以下的训练数据,Simplax 的优势恰恰在 17-30 提示的低条件区间最大,说明增益来自反向桥结构本身而非额外的条件拟合。

OpenWebText unconditional generation at selected NFE values. Gen. ENT is generative unigram entropy and should be compared with the data entropy 5.44. Gen. PPL is evaluated by the indicated external language model. The best and second-best values in each column are shown in bold and underlined, respectively.
Table 1: OpenWebText unconditional generation at selected NFE values. Gen. ENT is generative unigram entropy and should be compared with the data entropy 5.44. Gen. PPL is evaluated by the indicated external language model. The best and second-best values in each column are shown in bold and underlined, respectively.
Conditional Sudoku solving accuracy and unconditional Sudoku validity in percent. All models are trained with 30 clues. The 40- and 35-clue settings evaluate transfer to more heavily conditioned inputs, the 30-clue setting matches the training clue density, and the 25-, 20-, and 17-clue settings evaluate transfer to progressively less heavily conditioned inputs.
Table 2: Conditional Sudoku solving accuracy and unconditional Sudoku validity in percent. All models are trained with 30 clues. The 40- and 35-clue settings evaluate transfer to more heavily conditioned inputs, the 30-clue setting matches the training clue density, and the 25-, 20-, and 17-clue settings evaluate transfer to progressively less heavily conditioned inputs.
Design diagnostics on OpenWebText. The rows show temperature-swept generation frontiers at NFE = 16 and 128. The columns compare self-conditioning (a, d), denoiser input wt versus zt (b, e), and initialization from a pretrained UDLM checkpoint under a matched 1M-iteration budget (c, f). The dotted line marks the OpenWebText entropy, 5.44.
Figure 3: Design diagnostics on OpenWebText. The rows show temperature-swept generation frontiers at NFE = 16 and 128. The columns compare self-conditioning (a, d), denoiser input wt versus zt (b, e), and initialization from a pretrained UDLM checkpoint under a matched 1M-iteration budget (c, f). The dotted line marks the OpenWebText entropy, 5.44.
Llama-2 7B generative frontiers on OpenWebText for NFE ∈ {8, 16, 32, 64, 128, 256, 512, 1,024}. Each panel shows the temperature-swept Gen. PPL–Gen. ENT tradeoff. Lower Gen. PPL is better, and the reference data entropy is 5.44.
Figure 4: Llama-2 7B generative frontiers on OpenWebText for NFE ∈ {8, 16, 32, 64, 128, 256, 512, 1,024}. Each panel shows the temperature-swept Gen. PPL–Gen. ENT tradeoff. Lower Gen. PPL is better, and the reference data entropy is 5.44.

局限与改进

作者承认的局限:方法专为均匀范畴腐蚀定制,尚未推广到掩码/吸收核等更广的腐蚀核家族;辅助单纯形态相对标准离散扩散的计算开销没有被完整刻画;浓度调度 $\eta_t$ 是额外的设计选择(主实验取常数 0.01),理论并未给出其确定方式。我的补充观察:其一,数独优势随条件增强而收窄,40 提示下 98.55% 仅比 MDLM 的 98.45% 高 0.1 个百分点,收益主要来自低条件端(17-30 提示)与无条件端,暗示增益集中在需要逆向归纳的场景;其二,文本对比限于 OpenWebText 与 179M 规模,未涉及 MAJESTIC 等标准基准或更大模型,外推性未知;其三,评估依赖 GPT-2/Llama-2 的生成困惑度代理指标,与下游任务表现的关联未验证;其四,每步额外一次 Dirichlet 采样与 $w_s$ 维护的 wall-clock 成本没有数字支撑;其五,NFE=128 中档预算下 Llama-2 指标输给 LangFlow,并非所有工作点占优。

独立分析的弱点

弱点一:腐蚀核绑定均匀扩散。掩码扩散(MDLM、SEDD 一系)才是当前语言建模主流,而把 $w_t$ 的 shifted Dirichlet 锚定推广到吸收态并不显然——掩码态不属于词表单纯形的自然重心。改进方向:为吸收核设计对应增广,如把掩码态映到单纯形专属顶点并重推边际与桥。弱点二:$\eta_t$ 凭经验取常数 0.01,无论文级消融,作者也承认理论不决定它。改进方向:把 $\eta_t$ 与噪声调度 $\alpha_t$ 联动,补 sensitivity 曲线。弱点三:效率论证不完整,未报告训练/采样吞吐与显存对比,$w_s$ 维护成本未知。改进方向:给出与 UDLM 的逐步延迟、显存测量。弱点四:评估面窄,只有 OWT 无条件生成与数独,缺有条件生成与下游评测。改进方向:在 infilling、代码补全或 MAJESTIC 上复验。弱点五:17 提示下 1.20% 虽居首但绝对值接近失败,极低条件区间仍是开放难题,可结合约束解码或搜索式采样。

未来方向

作者明确提出的方向:把构造推广到更广的范畴腐蚀核;开发更高效的反向求解器。基于成果可以延伸的:其一,把 UDLM 预训练到 Simplax 微调的两阶段初始化(800k+200k 优于从零 1M 的前沿,Figure 3c,f)发展成通用的旧模型升级配方,让已发布的均匀扩散 checkpoint 免费换目标函数,这对实际部署很诱人;其二,用理论定 $\eta(t)$:连续时间极限 $\ell_{ct}$ 中 $\lambda(t)=-\frac{d}{dt}\log\alpha(t)$ 与浓度的耦合暗示可能存在使桥匹配最优的解析调度;其三,与均匀扩散已验证的少步生成、自我纠错、token 修订能力结合,在 NFE=16 一档的低预算场景做进一步加速;其四,迁移到生物序列(如 DDSM 验证过的 DNA 生成)等类别数大或结构约束强的范畴数据;其五,把辅助单纯形桥思想嫁接到掩码扩散或离散流匹配,检验它是否是跨腐蚀核通用的增强件;其六,给出 $w_t$ 信息量的理论刻画(它到底比 one-hot 锚点多利用了什么)。

复现评估

复现难度中等偏高。有利条件:正文加附录(A-C 证明、E.1 架构细节)把目标式、采样步骤与超参写得很具体;OpenWebText 公开,数独基准沿用 Deschenaux & Gulcehre (2026) 的 48,000/2,000 划分且协议清晰;主配置完整:GPT-2 BPE、L=1024、179M DiT、Adam lr $3\times10^{-4}$、batch 512、1M 迭代、$\eta\equiv0.01$,数独侧 20k 步、batch 256。不利条件:文本主实验是 512 批 x 1024 token x 1M 步量级,估计需多卡 A100/H100 数周,远超个人实验室预算;数独实验较小,单机可复现;论文文本未见官方代码发布说明,Dirichlet 采样、$\rho_{s|t}$ 的数值稳定性($\oslash$、$\odot$ 需防零除)与 15 档温度扫描评估管线都要自行实现;完整对比还需复现 CANDI、LangFlow 等 7 个基线。务实路径是先复现 50k 步设计诊断与数独实验验证实现正确性,再投入 1M 步主实验。