DGL 消息传递(Message Passing)完全指南:内置函数、高效实现与异构图 multi_update_all
2026/9/23 8:14:52 网站建设 项目流程

DGL 消息传递(Message Passing)完全指南:内置函数、高效实现与异构图 multi_update_all

【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl

本篇技术指南以 DGL 官方用户指南第二章 message.rst 及其四个子章节(message-api.rst、message-efficient.rst、message-part.rst、message-heterograph.rst)为主体,系统讲解 DGL 中消息传递的计算范式、dgl.function内置函数、apply_edges/update_all等核心 API、高效编码技巧,以及在子图和异构图上的应用。读完本文,你将能够用内置函数写出正确且高效的 GNN 消息传递代码,理解EdgeBatch/NodeBatch的底层数据结构,并掌握multi_update_all处理多关系图的方法。

消息传递范式:GNN 计算的基本抽象

在 DGL 中,图被定义为节点集合与边集合的拓扑结构,特征则挂在节点和边上。设节点特征为 $x_v \in \mathbb{R}^{d_1}$,边 $(u, v)$ 的特征为 $w_e \in \mathbb{R}^{d_2}$,消息传递范式在第 $t+1$ 步定义了两类计算:

边级计算(Edge-wise)——为每条边生成消息:

$$m_e^{(t+1)} = \phi \left( x_v^{(t)}, x_u^{(t)}, w_e^{(t)} \right), \quad (u, v, e) \in \mathcal{E}$$

节点级计算(Node-wise)——聚合入边消息并更新节点特征:

$$x_v^{(t+1)} = \psi \left(x_v^{(t)}, \rho\left(\left\lbrace m_e^{(t+1)} : (u, v, e) \in \mathcal{E} \right\rbrace \right) \right)$$

其中:

  • $\phi$ 是消息函数(message function),定义在每条边上,将边特征与其两端节点特征组合生成消息;
  • $\rho$ 是归约函数(reduce function),将节点收到的所有入边消息聚合(如summaxminmean);
  • $\psi$ 是更新函数(update function),定义在每个节点上,将聚合结果与节点自身特征结合并写回节点特征。

这一范式贯穿 DGL 全部消息传递 API。原章正文见 message.rst,其 Roadmap 将后续内容划分为四个主题:内置函数与 API、高效编码、子图上的消息传递、异构图消息传递,本文依次展开。

内置函数与消息传递 API

EdgeBatch 与 NodeBatch:UDF 的入参结构

在 DGL 中,消息函数接收唯一的参数edges,它是一个dgl.udf.EdgeBatch实例(定义见 python/dgl/udf.py)。消息传递过程中 DGL 内部生成该对象以表示一批边,它暴露三个成员:

  • edges.src:源节点特征视图;
  • edges.dst:目的节点特征视图;
  • edges.data:边特征视图。

此外还有edges.edges()返回边端点三元组(U, V, EID),以及edges.batch_size()返回批内边数。

归约函数接收唯一的参数nodes,它是一个dgl.udf.NodeBatch实例(python/dgl/udf.py),其成员mailbox用于访问该批节点收到的消息。mailbox['m']的形状为(N, D, ...),其中N是本批节点数、D是每个节点收到的消息数,因此对消息求和时需对dim=1归约。NodeBatch还提供nodes.data(节点特征)、nodes.nodes()(节点 ID)和nodes.batch_size()

更新函数同样接收nodes参数,作用于归约函数的结果,通常将其与节点原始特征组合,并把结果保存为节点特征。

优先使用 dgl.function 内置函数

DGL 在命名空间dgl.function(即fn)下实现了常用消息函数与归约函数的内置版本(built-in)。DGL 官方建议只要可能就使用内置函数,因为它们经过深度优化,并且自动处理维度广播(broadcasting)。

内置消息函数分为一元与二元两类:

  • 一元(unary):支持copy,例如copy_ucopy_e
  • 二元(binary):支持addsubmuldivdot

命名约定为:u代表源节点(src),v代表目的节点(dst),e代表边(edge)。参数均为字符串,分别指定输入输出字段名。例如把源节点hu特征与目的节点hv特征相加、结果保存到边的he字段,可写:

import dgl.function as fn fn.u_add_v('hu', 'hv', 'he')

它等价于下面的消息 UDF:

def message_func(edges): return {'he': edges.src['hu'] + edges.dst['hv']}

内置归约函数支持summaxminmean(源码见 python/dgl/function/reducer.py,通过_gen_reduce_builtin动态生成并注册)。归约函数通常有两个字符串参数:mailbox中的消息字段名与节点特征字段名。例如dgl.function.sum('m', 'h')等价于:

import torch def reduce_func(nodes): return {'h': torch.sum(nodes.mailbox['m'], dim=1)}

二元消息函数的完整集合在 python/dgl/function/message.py 中通过_register_builtin_message_func动态生成:对目标组合u/v/e两两配对(lhs != rhs)逐一注册add/sub/mul/div/dot五种运算,因此实际可用函数包括u_add_vu_mul_ev_dot_ee_sub_u等 30 个组合;一元复制函数copy_ucopy_e则显式定义于 python/dgl/function/message.py。当内置函数无法表达需求时,再实现用户自定义的 message/reduce 函数(UDF)。

apply_edges:仅做边级计算

apply_edges只调用边级计算、不触发消息传递,接收一个消息函数为参数,默认更新所有边的特征(python/dgl/heterograph.py)。它也支持通过edges参数限定要更新的边(边 ID、节点对张量等形式),并可通过etype指定边类型。例如:

import dgl.function as fn graph.apply_edges(fn.u_add_v('el', 'er', 'e'))

update_all:消息传递一站式 API

update_all是高层次 API,将消息生成、消息聚合、节点更新合并为一次调用,从而为整体优化(如内存复用)留出空间。其签名为(python/dgl/heterograph.py):

update_all(message_func, reduce_func, apply_node_func=None, etype=None)

三个核心参数分别为消息函数、归约函数与更新函数,etype用于异构图中指定边类型。DGL 推荐把更新函数放到update_all之外、不作为参数传入,因为更新函数通常可以用纯张量运算简洁表达。例如:

import dgl.function as fn def update_all_example(graph): # 结果保存在 graph.ndata['ft'] graph.update_all(fn.u_mul_e('ft', 'a', 'm'), fn.sum('m', 'ft')) # 在 update_all 之外调用更新函数 final_ft = graph.ndata['ft'] * 2 return final_ft

该调用将源节点特征ft与边特征a相乘生成消息m,将消息m求和更新节点特征ft,最后将ft乘以 2 得到final_ft。调用结束后,DGL 会清理中间消息m。上述代码的数学表达式为:

$$final_ft_i = 2 \times \sum_{j \in \mathcal{N}(i)} (ft_j \times a_{ji})$$

浮点类型支持与 float16

DGL 内置函数支持浮点数据类型,即特征必须是halffloat16)/float/double张量。其中float16支持默认关闭,因为它对 GPU 有最低算力要求:计算能力需不低于sm_53(即 Pascal、Volta、Turing 和 Ampere 架构)。如需为混合精度训练启用 float16,需要从源码编译 DGL,具体步骤参见 Mixed Precision Training 教程。

编写高效的消息传递代码

DGL 对消息传递的内存消耗与计算速度做了专门优化。利用这些优化的常见做法是:用内置函数作为参数,将自定义消息传递逻辑组织成若干次update_all调用的组合

避免从节点到边的多余内存拷贝

对于某些图,边的数量远大于节点数量,此时应尽量避免把节点特征拷贝到边上。但有些场景(如dgl.nn.pytorch.conv.GATConv,GAT 需要把消息保存在边上用于后续 softmax 等操作)必须调用apply_edges配合内置函数在边上保存消息。由于边上的消息可能是高维的、非常耗内存,DGL 建议尽可能保持边特征维度尽量低

下面是一个把边上的运算拆分到节点上执行的经典例子。目标是拼接源特征与目的特征再经过线性层,即 $W \times (u \Vert v)$,其中srcdst特征维度高,而线性层输出维度低。

直接实现(低效)——先拼接到边上再乘线性层:

import torch import torch.nn as nn linear = nn.Parameter(torch.FloatTensor(size=(node_feat_dim * 2, out_dim))) def concat_message_function(edges): return {'cat_feat': torch.cat([edges.src['feat'], edges.dst['feat']], dim=1)} g.apply_edges(concat_message_function) g.edata['out'] = g.edata['cat_feat'] @ linear

推荐实现(高效)——利用等式 $W \times (u \Vert v) = W_l \times u + W_r \times v$($W_l$、$W_r$ 分别是矩阵 $W$ 的左半与右半),把线性层拆成两个,分别作用在源特征与目的特征上,最后在边上相加:

import dgl.function as fn linear_src = nn.Parameter(torch.FloatTensor(size=(node_feat_dim, out_dim))) linear_dst = nn.Parameter(torch.FloatTensor(size=(node_feat_dim, out_dim))) out_src = g.ndata['feat'] @ linear_src out_dst = g.ndata['feat'] @ linear_dst g.srcdata.update({'out_src': out_src}) g.dstdata.update({'out_dst': out_dst}) g.apply_edges(fn.u_add_v('out_src', 'out_dst', 'out'))

两种实现在数学上等价。后者更高效的原因在于:不需要把feat_srcfeat_dst保存在边上(省内存),且加法可以用 DGL 内置函数fn.u_add_v完成,进一步加速计算、缩减内存占用。完整说明见 message-efficient.rst。

在图的子图上应用消息传递

如果只想更新图中的部分节点,标准做法是先用节点 ID 构造子图,再在子图上调用update_all

nid = [0, 2, 3, 6, 7, 9] sg = g.subgraph(nid) sg.update_all(message_func, reduce_func, apply_node_func)

这是小批量(mini-batch)训练中的常见用法,例如邻居采样后对采样得到的子图执行消息传递,避免在整个大图上计算。更详细的用法参见 minibatch.rst(对应guide-minibatch章节)。此处完整继承自 message-part.rst。

异构图上的消息传递

异构图(heterogeneous graph,简称 heterograph)包含不同类型的节点与边,不同类型的节点和边往往拥有不同类型的属性,用于刻画各自的特性(异构图的构建与表示参见 graph-heterogeneous.rst)。在图神经网络语境下,根据复杂度不同,某些节点类型和边类型可能需要用不同维数的表示来建模。

异构图上消息传递可拆为两步:

  1. 对每个关系 r 分别做消息计算与聚合
  2. 归约(reduction):把每个节点类型在所有关系上的聚合结果合并。

DGL 在异构图上调用消息传递的接口是multi_update_all(python/dgl/heterograph.py)。它接收两个参数:

  • 字典:以关系(relation)为键,值为该关系下update_all的参数((message_func, reduce_func, [apply_node_func]));
  • 字符串:跨类型归约器(cross type reducer),可取summinmaxmeanstack

一个典型示例(R-GCN 风格的多关系消息传递):

import dgl.function as fn for c_etype in G.canonical_etypes: srctype, etype, dsttype = c_etype Wh = self.weightetype # 将变换结果保存在图中供消息传递使用 G.nodes[srctype].data['Wh_%s' % etype] = Wh # 为每个关系指定消息传递函数: (message_func, reduce_func) # 注意结果都保存到同一个目的特征 'h',这提示了按类型归约的方式 funcs[etype] = (fn.copy_u('Wh_%s' % etype, 'm'), fn.mean('m', 'h')) # 触发多类型消息传递 G.multi_update_all(funcs, 'sum') # 返回更新后的节点特征字典 return {ntype: G.nodes[ntype].data['h'] for ntype in G.ntypes}

其中G.canonical_etypes给出形如(srctype, etype, dsttype)的规范边类型三元组;每个关系先用对应关系的权重矩阵self.weight[etype]对源节点特征做线性变换,保存为Wh_<etype>;再用fn.copy_u复制为消息m、以fn.mean归约到目的特征h;最后用multi_update_all(funcs, 'sum')对所有关系的聚合结果做跨类型求和。由于各关系的结果写入了同一个目的特征h,跨类型归约器才能正确合并它们。该示例完整继承自 message-heterograph.rst。

小结与进一步阅读

围绕消息传递,DGL 提供了从EdgeBatch/NodeBatchUDF、dgl.function内置函数、apply_edges/update_all高层 API,到子图更新与multi_update_all异构图接口的完整体系。实践要点可归纳为:

  1. 优先内置函数:性能优且自动广播,UDF 仅在内置无法表达时使用;
  2. 更新函数外置:将update_all外的更新写成纯张量运算,代码更简洁;
  3. 降低边特征维度:把拼接等重操作拆分到节点上执行,减少节点→边的内存拷贝;
  4. 子图 + update_all:小批量训练的标准消息传递模式;
  5. multi_update_all:异构图按关系传递消息后跨类型归约。

如需继续深入,可阅读源码 python/dgl/function/message.py 与 python/dgl/function/reducer.py 了解内置函数生成机制,或阅读 python/dgl/udf.py 掌握 UDF 的完整 API(包括edges()batch_size()等高级用法),亦可在 DGL 官方 API 参考文档中查看全部内置函数列表。

【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询