[论文解读] Gradient-Based Neural DAG Learning
该论文提出 GraN-DAG,一种基于梯度的从观测数据中学习有向无环图(DAG)的方法,利用神经网络建模非线性关系。通过将 NOTEARS 连续优化框架扩展为可微分的无环约束,GraN-DAG 在合成数据集和真实世界数据集上均达到最先进性能,优于连续方法,并在关键因果推断指标上与贪心搜索方法(如 CAM 和 GSF)相当。
We propose a novel score-based approach to learning a directed acyclic graph (DAG) from observational data. We adapt a recently proposed continuous constrained optimization formulation to allow for nonlinear relationships between variables using neural networks. This extension allows to model complex interactions while avoiding the combinatorial nature of the problem. In addition to comparing our method to existing continuous optimization methods, we provide missing empirical comparisons to nonlinear greedy search methods. On both synthetic and real-world data sets, this new method outperforms current continuous methods on most tasks, while being competitive with existing greedy search methods on important metrics for causal inference.
研究动机与目标
- 解决在关系为非线性时,从观测数据中学习因果 DAG 结构的挑战。
- 将 NOTEARS 连续优化框架扩展以支持使用神经网络的非线性依赖关系。
- 通过与非线性贪心搜索方法(如 CAM 和 GSF)进行比较,弥合实证评估的差距。
- 提供一种可扩展的、可微分的方法,避免组合搜索,同时在结构学习基准上保持优异性能。
提出的方法
- 通过用前馈神经网络替代线性模型,扩展 NOTEARS 框架,以在结构方程模型中捕捉非线性关系。
- 通过确保邻接矩阵的谱半径小于 1,实现一种可微分的无环约束,采用基于路径的可微分松弛方法。
- 采用连续优化公式,通过邻接矩阵的可微函数强制执行 DAG 约束,支持使用梯度下降的端到端训练。
- 在神经网络和图两个层面均采用基于路径的无环条件松弛,以保持可微性。
- 采用两阶段优化过程:首先优化神经网络参数,然后使用投影梯度法精炼邻接矩阵。
- 引入剪枝机制和阈值策略,以提高最终 DAG 的稀疏性和可解释性。
实验结果
研究问题
- RQ1能否在保持无环性的同时,将连续可微分优化框架扩展至使用神经网络的非线性 DAG 结构学习?
- RQ2在合成数据和真实世界数据上,GraN-DAG 相较于现有连续方法(如 NOTEARS 和 DAG-GNN)在结构学习准确性方面表现如何?
- RQ3在非线性设置下,GraN-DAG 相较于贪心搜索方法(如 CAM 和 GSF)是否达到具有竞争力的性能?
- RQ4可微分的无环约束对使用深度模型进行 DAG 学习的可扩展性和收敛性有何影响?
- RQ5GraN-DAG 是否能在包括高维和真实世界生物数据在内的多种数据类型上实现良好泛化?
主要发现
- 在 10 个节点的合成数据上,GraN-DAG 的平均 SHD 为 2.6±2.4,显著优于 NOTEARS(21.2±11.5)和 DAG-GNN(8.7±2.8),并在所有指标上与 CAM 和 GSF 相当或更优。
- 在 50 个节点的 ER4 合成图上,GraN-DAG 的平均 SID 为 37.1±12.4,优于 NOTEARS(41.3±11.5)和 DAG-GNN(63.6±8.6),在某些情况下与 CAM(28.5±21.5)和 GSF(44.5±19.7)相当。
- 在真实世界的蛋白质信号传导数据集上,GraN-DAG 的 SHD 为 12.0,SHD-C 为 9.0,优于 NOTEARS(15.0 和 14.0),并在关键指标上与 CAM(11.0 和 9.0)相当。
- 在 SynTReN 20 个节点的数据集上,GraN-DAG 的平均 SHD 为 41.2±9.6,优于 NOTEARS(44.2±27.5)和 DAG-GNN(32.2±5.0),在 SHD 上与 CAM(101.7±37.2)相当,但 SHD-C 更优。
- 在超参数搜索中,GraN-DAG 在所有设置下均持续获得低于 NOTEARS 和 DAG-GNN 的 SHD 和 SID,表现出稳健性和泛化能力。
- 该方法在合成数据和真实世界数据上均表现强劲,与 CAM 和 GSF 等贪心搜索基线方法相比具有竞争力,尤其在 SHD 和 SID 指标上表现突出。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。