发表机构
Tsinghua University; Beijing Institute of Mathematical Sciences and Applications; Wuhan University; MathonAI(清华大学; 北京数学科学与应用研究院; 武汉大学; 马松人工智能)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
研究针对分数熵离散扩散存在的问题,引入均值到分数(M2S)方法,通过预测后验均值并转换为分数,应用于特定马尔可夫链。在CIFAR-10等实验中,M2S降低了测试BPD和FID-50k,提升了生成性能。
AI 中文摘要
分数熵离散扩散(SEDD)使用无约束的正分数比来参数化离散反向过程。虽然正性保证了非负的反向跳跃率,但它并不能确保贝叶斯可实现性:噪声状态下的比率不一定由前向核下的任何干净令牌后验联合诱导。分数熵损失有正确的总体最优值,但不能在远离最优值时强制执行此约束。在训练好的纯均匀SEDD检查点中,大约四分之一的完整分数向量违反坐标框,而超过一半位于框内但仍与任何有效后验在实质上不兼容。这种违反会在有限步采样中产生负的预归一化权重。将原始分数投影到桥多面体上可以消除所有观察到的负权重,并在不改变采样器的情况下将外部生成PPL从203.6提高到175.1。我们引入了均值到分数(M2S),它预测干净令牌后验均值并通过精确的核相关线性映射将其转换为分数。该构造适用于任何满足温和支持条件的已知坐标连续时间马尔可夫链(CTMC)。对于均匀损坏,它将概率单纯形映射到桥多面体上;对于吸收掩码损坏,得到的目标精确地恢复了MD4。在一个受控的2840万参数CIFAR-10比较中,M2S将测试BPD从3.173降低到3.129,将FID-50k从CifarSEDDFID降低到CifarMtwoSFID。一个在大约262B OpenWebText令牌插槽上训练的1.7亿参数M2S模型在每个测试采样预算下都优于评估的纯均匀SEDD、GIDD和神经CTMC检查点,在128步时达到生成PPL 143.3,而最强的纯均匀基线为183.6。
英文摘要
Score Entropy Discrete Diffusion (SEDD) parameterizes discrete reverse processes with unconstrained positive score ratios. While positivity guarantees nonnegative reverse jump rates, it does not ensure Bayes realizability: ratios at a noisy state need not be jointly induced by any clean-token posterior under the forward kernel. The score-entropy loss has the correct population optimum but does not enforce this constraint away from it. In a trained pure-uniform SEDD checkpoint, roughly one quarter of complete score vectors violate the coordinate box, while more than half lie inside it yet remain materially incompatible with any valid posterior. Such violations can produce negative pre-normalization weights in finite-step sampling. Projecting raw scores onto the bridge polytope removes all observed negative weights and improves external generative PPL from $203.6$ to $175.1$ without changing the sampler. We introduce \emph{mean-to-score} (M2S), which predicts a clean-token posterior mean and converts it to the score through an exact kernel-dependent linear map. The construction applies to any known coordinate-wise continuous-time Markov chain (CTMC) satisfying a mild support condition. For uniform corruption, it maps the probability simplex onto the bridge polytope; for absorbing-mask corruption, the resulting objective recovers MD4 exactly. In a controlled 28.4M-parameter CIFAR-10 comparison, M2S lowers test BPD from $3.173$ to $3.129$ and FID-50k from $\CifarSEDDFID$ to $\CifarMtwoSFID$. A 170M-parameter M2S model trained on about 262B OpenWebText token slots outperforms the evaluated pure-uniform SEDD, GIDD, and Neural CTMC checkpoints at every tested sampling budget, reaching generative PPL $143.3$ at 128 steps versus $183.6$ for the strongest pure-uniform baseline.