arXiv:2411.16142cs.LGstat.ML2024-11被引 2

通过因果学习提升图上时空预测的分布外泛化能力

Causal Adjacency Learning for Spatiotemporal Prediction Over Graphs

  • 基于因果推断构建图结构的邻接矩阵,避免依赖固定距离或相关性
  • 在分布外测试数据上显著提升预测精度,验证了方法鲁棒性
  • 适合关注模型泛化能力的交通预测研究者

图上时空预测(STPG)对交通系统至关重要。现有模型多通过距离或相关性直接构造邻接矩阵,但这类方法未考虑测试数据潜在的模式变化,在分布外(OOD)场景下性能下降。本文提出因果邻接学习(CAL)方法,挖掘图上节点间的因果关系。所学因果邻接矩阵在真实世界图数据上的下游预测任务中表现优异,即使下游任务未显式进行因果建模,也能有效提升对分布外数据的预测性能。

原文摘要 · Abstract (English)

Spatiotemporal prediction over graphs (STPG) is crucial for transportation systems. In existing STPG models, an adjacency matrix is an important component that captures the relations among nodes over graphs. However, most studies calculate the adjacency matrix by directly memorizing the data, such as distance- and correlation-based matrices. These adjacency matrices do not consider potential pattern shift for the test data, and may result in suboptimal performance if the test data has a different distribution from the training one. This issue is known as the Out-of-Distribution generalization problem. To address this issue, in this paper we propose a Causal Adjacency Learning (CAL) method to discover causal relations over graphs. The learned causal adjacency matrix is evaluated on a downstream spatiotemporal prediction task using real-world graph data. Results demonstrate that our proposed adjacency matrix can capture the causal relations, and using our learned adjacency matrix can enhance prediction performance on the OOD test data, even though causal learning is not conducted in the downstream task.

时空预测因果学习图神经网络分布外泛化

Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。