中文

一种用于图神经网络学习的拟 Wasserstein 损失

机器学习 2024-03-14 v4

摘要

在节点级预测任务中学习图神经网络(GNNs)时,大多数现有损失函数对每个节点独立施加,尽管由于图结构,节点嵌入及其标签是非独立同分布(non-i.i.d.)的。为消除这种不一致性,本研究借助图上定义的最优传输,提出一种新颖的拟 Wasserstein(QW)损失,从而引出 GNNs 新的学习与预测范式。具体而言,我们设计了观测的多维节点标签与其估计之间的“拟 Wasserstein”距离,优化定义于图边上的标签传输。估计由 GNN 参数化,其中最优标签传输可选择性地决定图边权重。通过将标签传输的严格约束重构为基于 Bregman 散度的正则项,我们得到所提的拟 Wasserstein 损失,并配有联合学习 GNN 与最优标签传输的两个高效求解器。在预测节点标签时,我们的模型将 GNN 的输出与最优标签传输提供的残差分量相结合,形成一种新颖的转导预测范式。实验表明,所提 QW 损失适用于多种 GNNs,并有助于提升其在节点级分类与回归任务中的性能。本工作代码见 \url{https://github.com/SDS-Lab/QW_Loss}。

关键词

引用

@article{arxiv.2310.11762,
  title  = {A Quasi-Wasserstein Loss for Learning Graph Neural Networks},
  author = {Minjie Cheng and Hongteng Xu},
  journal= {arXiv preprint arXiv:2310.11762},
  year   = {2024}
}