Neural Estimation of Pairwise Mutual Information in Masked Discrete Sequence Models
原作者: Jai Sharma, Yifan Wang, Bryan Li
原作者: Jai Sharma, Yifan Wang, Bryan Li
原始论文采用 CC BY 4.0 许可(http://creativecommons.org/licenses/by/4.0/)。 ✨ 这是对下方论文的AI生成解释。它不是由作者撰写或认可的。如需技术准确性,请参阅原始论文。 阅读完整免责声明
技术摘要:掩码离散序列模型中的成对互信息神经估计
1. 问题陈述
掩码扩散模型(MDMs)是用于离散序列(如文本、蛋白质、数独)的强大生成模型,它们避免了自回归(AR)模型中固定的回归排序。然而,标准的 MDM 主要暴露边际条件分布(p(xi∣xcontext)),并未显式表示变量间依赖关系。
这种缺乏显式依赖建模的情况带来了两个主要挑战:
- 可解释性:难以理解模型内部关于变量如何相互关联的信念结构。
- 并行解码效率:当前的并行解码策略(如 Mask-Predict、EB-Sampler)通常依赖边际置信度(熵)来决定同时解除哪些标记的掩码。这种方法未能考虑成对依赖关系。在不相互条件的情况下同时解除高度相关的标记(高互信息),会导致全局不一致性(例如违反数独规则或蛋白质结构约束),往往迫使回退到顺序解码,或导致生成质量低下。
由于需要进行密度估计,传统互信息(MI)计算在高维设置中在计算上是不可行的。
2. 方法论
作者提出了一种神经框架,直接从预训练 MDM 的隐藏状态中估计成对条件互信息(I(Xi;Xj∣C))。该方法包含三个主要组成部分:
A. 真实 MI 计算(监督信号)
为了训练一个轻量级估计器,作者首先定义了一种精确但昂贵的基于预训练 MDM 自身条件分布来计算“真实”MI 的方法。
- 定义:对于上下文 C(未掩码的标记),两个掩码位置 i 和 j 之间的 MI 定义为联合分布 P(Xi,Xj∣C) 与边际分布乘积之间的 KL 散度。
- 计算策略:由于 MDM 输出边际分布,作者使用了一种基于扰动的暴力探测策略:
- 基础传递:在掩码序列上运行模型以获得边际分布 P(Xi∣C) 并计算个体熵 H(Xi∣C)。
- 条件传递:对于每个位置 i 和每个可能的标记 v,固定 Xi=v 并运行前向传递以获得条件分布 P(Xj∣Xi=v,C)。
- 计算:计算条件熵 H(Xj∣Xi,C) 并将 MI 推导为熵的减少量:I(Xi;Xj∣C)=H(Xj∣C)−H(Xj∣Xi,C)。
- 成本:这需要 1+N⋅∣V∣ 次前向传递,使其在推理阶段不可行,但适合生成训练数据。
B. 神经 MI 估计器
一个轻量级神经网络(fϕ)被训练以直接从冻结的 MDM 隐藏状态(h)中近似 MI 矩阵。
- 架构:估计器接收隐藏状态 h∈RN×D 并输出一个对称矩阵 I^∈RN×N,表示所有位置的估计成对 MI。
- 训练目标:模型被训练以最小化预测矩阵 I^ 与掩码索引上的真实矩阵 MGT 之间的均方误差(MSE)。
C. MI 引导的并行采样
作者引入了一种用于并行解码的贪婪选择算法,该算法利用预测的 MI 矩阵来确保未掩码标记之间的条件独立性。
- 策略:算法不是简单地选择熵最低(置信度最高)的标记,而是选择一个标记批次 S,使得它们在给定上下文的情况下相互独立。
- 算法:
- 按熵递增(即置信度最高优先)对掩码索引进行排序。
- 遍历候选项,计算依赖成本:d(i∣U)=∑j∈UI^i,j,其中 U 是已选标记的集合。
- 仅当标记 i 的总成本(熵 + λ× 依赖成本)在剩余预算 γ 内时,才选择该标记。
- 如果成本过高(表明与已选标记具有高 MI),则该标记被推迟到顺序步骤中处理。
- 结果:这确保了高度相关的变量按顺序处理,而条件独立的子集则并行处理。
3. 主要贡献
- 神经 MI 估计框架:一种直接从 MDM 隐藏状态估计成对条件 MI 的方法,绕过了推理过程中对昂贵密度估计或真实值计算的需求。
- MI 引导的并行解码:一种新颖的采样策略,利用估计的 MI 来识别变量的条件独立子集,从而实现既能保持全局一致性又能安全并行化的解码。
- 可解释性工具:MI 图作为模型内部信念结构的可视化工具,揭示了无需显式编程即可学到的约束(例如数独规则、蛋白质折叠依赖关系)。
4. 实验结果
该方法在两个领域进行了评估:数独(结构化逻辑)和蛋白质序列生成(使用 ESM-C)。
数独
- 设置:在 100,000 个谜题上训练;在 1,000 个未见过的困难谜题上评估。
- 性能:
- 顺序基线:平均 53.9 次前向传递,准确率 61.6%。
- 朴素并行(k=7):9.0 次传递,但准确率降至 36.8%。
- MI 引导(γ=0.3):15.2 次传递,准确率为63.6%(超过顺序基线)。
- MI 引导(γ=0.6):9.7 次传递,准确率为 56.2%。
- 观察:与顺序解码相比,MI 引导采样器将前向传递次数减少了 3-5 倍,同时与朴素并行方法相比保持或提高了准确率。
蛋白质序列(ESM-C)
- 设置:生成 500 个随机蛋白质(长度 50-100),并使用 Jensen-Shannon 散度(JSD)与来自 UniRef50 的 500 个参考样本进行比较。
- 性能:
- 顺序:74.8 次传递,JSD 为 0.093。
- 朴素并行(k=12):6.2 次传递,JSD 为 0.218(质量显著下降)。
- MI 引导(γ=4):10.0 次传递,JSD 为 0.174。
- 观察:MI 引导采样实现了比朴素并行基线更好的速度 - 准确率权衡,显著减少了传递次数(与顺序解码相比接近一个数量级),同时比基于熵的方法更好地保持了生成质量。
5. 意义与主张
该论文声称,显式建模变量依赖关系对于释放离散扩散模型的全部潜力至关重要。
- 弥合差距:这项工作弥合了顺序采样的高质量与并行解码的高效性之间的差距。
- 内部表示:MI 图表明,MDM 在没有显式编程的情况下自然地获得了刚性结构约束(如数独规则或蛋白质依赖关系),并且可以通过估计器提取这些约束。
- 效率:该方法实现了 MI 引导的并行解码,能够识别条件独立子集,从而与顺序解码相比,将推理时的前向传递次数减少了 3-5 个数量级。
承认的局限性:
作者指出,预测器并不完美,需要大量的设置和训练(在训练数据上即时计算真实 MI)。未来的工作建议研究最佳的预测器架构和改进的课程训练策略,以避免在训练阶段计算真实值的计算成本。
您所在领域的论文太多了?
获取与您研究关键词匹配的最新论文每日摘要——附技术摘要,使用您的语言。
每周获取最佳 machine learning 论文。
受到斯坦福、剑桥和法国科学院研究人员的信赖。
请查收邮箱确认订阅。
出了点问题,再试一次?
无垃圾邮件,随时退订。