[论文解读] Graph Representation Learning via Multi-task Knowledge Distillation
本文提出了一种多任务知识蒸馏框架,通过利用基于网络理论的图度量(如密度和直径)作为辅助任务,提升图表示学习性能。通过联合训练主图级别预测任务与这些辅助度量,该方法在低标签设置下显著提升了模型表现,尤其在合成数据集和真实世界数据集(如NCI1和IMDB-BINARY)上展现出一致的性能增益。
Machine learning on graph structured data has attracted much research interest due to its ubiquity in real world data. However, how to efficiently represent graph data in a general way is still an open problem. Traditional methods use handcraft graph features in a tabular form but suffer from the defects of domain expertise requirement and information loss. Graph representation learning overcomes these defects by automatically learning the continuous representations from graph structures, but they require abundant training labels, which are often hard to fulfill for graph-level prediction problems. In this work, we demonstrate that, if available, the domain expertise used for designing handcraft graph features can improve the graph-level representation learning when training labels are scarce. Specifically, we proposed a multi-task knowledge distillation method. By incorporating network-theory-based graph metrics as auxiliary tasks, we show on both synthetic and real datasets that the proposed multi-task learning method can improve the prediction performance of the original learning task, especially when the training data size is small.
研究动机与目标
- 解决在标注数据稀缺时图级别表示学习性能不佳的挑战。
- 将网络理论中的领域知识整合到深度图表示模型中,以减少对大规模标注数据集的依赖。
- 通过多任务知识蒸馏提升图表示学习的泛化能力和样本效率。
- 证明基于图度量(如直径、密度)的辅助任务可有效将知识迁移至主预测任务。
- 在低数据设置下,于合成数据集和真实世界图基准数据集上验证该方法的有效性。
提出的方法
- 该方法采用共享的图表示主干网络(受DeepGraph启发),利用热核签名(HKS)从原始图中学习连续的、数据驱动的表示。
- 引入基于网络理论度量的多个辅助任务,具体为直接从图结构计算出的密度和直径。
- 构建多任务学习目标,联合最小化主任务的预测损失与辅助任务的损失,并引入可学习的任务权重(αk)以平衡各任务的贡献。
- 通过将图度量作为软标签应用知识蒸馏,使模型在推理阶段无需重新计算即可学习到结构归纳偏置。
- 模型架构在所有任务间共享前两个神经网络模块(卷积层和首个全连接层),而每个任务拥有独立的最终全连接层用于预测。
- 该方法可兼容任何监督式图表示学习模型,作为插件模块集成于现有框架之上。
实验结果
研究问题
- RQ1将基于网络理论的图度量作为辅助任务是否能提升图级别表示学习模型的性能?
- RQ2在标注数据有限的情况下,从手工设计的图度量中进行知识蒸馏是否能增强模型的泛化能力?
- RQ3辅助任务带来的性能增益如何随训练数据量变化,尤其是在低数据设置下?
- RQ4所提出的方法在合成数据集和真实世界图数据集上是否均有效?
- RQ5使用从图度量导出的软标签是否能相比实时计算特征降低推理成本?
主要发现
- 在合成的泊松随机图和优先连接图上,多任务模型始终优于单任务模型,且在训练数据量较小时性能差距最大。
- 在NCI1数据集上,多任务模型在10折交叉验证中的准确率高于单任务模型,尤其在仅使用少量训练数据时表现更优。
- 在IMDB-BINARY数据集上,多任务模型也表现出性能提升,但由于数据集本身较小导致方差较高,差异不如其他数据集明显。
- 随着训练数据量的增加,多任务与单任务模型之间的性能差距逐渐缩小,表明辅助任务在标签稀缺时最具优势。
- 该方法通过在训练阶段将知识蒸馏进模型,避免了推理时对图度量的实时计算,从而降低了测试时间的计算成本。
- 将图度量作为辅助任务可提升模型泛化能力,尤其在具备领域知识但标注数据有限的情况下效果更显著。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。