Skip to main content
QUICK REVIEW

[论文解读] Truncated Matrix Power Iteration for Differentiable DAG Learning

Zhen Zhang, Ignavier Ng|arXiv (Cornell University)|Aug 30, 2022
Remote-Sensing Image Classification被引用 4
一句话总结

该论文提出了一种新颖的可微分DAG学习方法,利用截断矩阵幂迭代(TMPI)来近似基于几何级数的DAG约束,从而在不引起数值不稳定性的情况下,对高阶多项式项施加更大的系数。该方法在结构汉明距离(SHD)上相比最先进方法最高提升3倍,尤其在稀疏图中表现更优,通过缓解梯度消失问题同时保持计算效率。

ABSTRACT

Recovering underlying Directed Acyclic Graph (DAG) structures from observational data is highly challenging due to the combinatorial nature of the DAG-constrained optimization problem. Recently, DAG learning has been cast as a continuous optimization problem by characterizing the DAG constraint as a smooth equality one, generally based on polynomials over adjacency matrices. Existing methods place very small coefficients on high-order polynomial terms for stabilization, since they argue that large coefficients on the higher-order terms are harmful due to numeric exploding. On the contrary, we discover that large coefficients on higher-order terms are beneficial for DAG learning, when the spectral radiuses of the adjacency matrices are small, and that larger coefficients for higher-order terms can approximate the DAG constraints much better than the small counterparts. Based on this, we propose a novel DAG learning method with efficient truncated matrix power iteration to approximate geometric series based DAG constraints. Empirically, our DAG learning method outperforms the previous state-of-the-arts in various settings, often by a factor of $3$ or more in terms of structural Hamming distance.

研究动机与目标

  • 解决由于高阶多项式项系数过小导致的可微分DAG学习中的梯度消失问题。
  • 证明当邻接矩阵的谱半径较小时,对高阶项使用更大系数是安全且有益的。
  • 开发一种高效算法,以有界误差和低计算成本近似基于几何级数的DAG约束。
  • 在合成与真实世界设置中,提升DAG学习的准确性和鲁棒性。
  • 将现有最先进模型中的DAG约束替换为基于TMPI的约束,以获得性能提升。

提出的方法

  • 提出一种基于几何级数的DAG约束,作为邻接矩阵上的d阶多项式,通过在高阶项使用更大系数,以更好地逼近幂零性条件。
  • 引入截断矩阵幂迭代(TMPI),一种高效算法,可在O(log k)时间内计算几何级数近似,其中k为级数的有效阶数。
  • TMPI算法保持理论误差界,确保近似值始终在真实几何级数的可控容差范围内。
  • 该方法利用DAG邻接矩阵具有幂零性这一事实,即其谱半径较小,从而即使在大系数下也不会引发数值爆炸。
  • 通过将该约束替换现有可微分DAG框架(如NOTEARS、DAG-GNN和GRAN-DAG)中的原始多项式无环约束,实现其集成。
  • 采用启发式方法识别一个约化阶数k ≤ d,使得约束的可行集保持不变,从而实现进一步的计算节省。

实验结果

研究问题

  • RQ1在不引起数值不稳定性的情况下,是否可以通过增大高阶多项式项的系数来提升DAG约束的准确性?
  • RQ2是否可以使用截断迭代方法,以有界误差高效近似邻接矩阵的几何级数?
  • RQ3用基于几何级数的约束替代现有DAG约束,是否能在多种图结构中带来SHD的显著改进?
  • RQ4所提出的TMPI算法在速度和准确性上与朴素实现及现有几何级数约束实现相比如何?
  • RQ5该方法是否能有效缓解稀疏DAG中梯度消失问题,其中高阶项对强制无环性至关重要?

主要发现

  • 所提出的基于TMPI的DAG约束在各种设置下,将结构汉明距离(SHD)降低了最多3倍,优于最先进方法。
  • 在50个节点的ER1非线性SEM中,该方法实现了22.2±4.2的SHD,优于DAG-GNN的25.2±4.5。
  • 在50个节点的非线性MLP数据集中,使用TMPI约束的NOTEARS-MLP实现了14.9±1.3的SHD,优于原始NOTEARS-MLP的16.9±1.5。
  • 在Sachs蛋白信号传导数据集中,将DAG约束替换为TMPI后,SHD从16降至16(DAG-GNN),从13降至12(Gran-DAG),对应SHDC提升从21降至17,从11降至9。
  • 快速的TMPI实现相比朴素实现显著更快,尤其在大图上,同时保持相近的SHD性能。
  • 该方法通过使用更大系数使高阶项具有信息量,有效缓解了稀疏图中的梯度消失问题,且未引发梯度或数值爆炸。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。