ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

DGL 消息传递(Message Passing)完全指南:内置函数、高效实现与异构图 multi_update_all 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}$消息传递范式在第 $t1$ 步定义了两类计算边级计算Edge-wise——为每条边生成消息$$m_e^{(t1)} \phi \left( x_v^{(t)}, x_u^{(t)}, w_e^{(t)} \right), \quad (u, v, e) \in \mathcal{E}$$节点级计算Node-wise——聚合入边消息并更新节点特征$$x_v^{(t1)} \psi \left(x_v^{(t)}, \rho\left(\left\lbrace m_e^{(t1)} : (u, v, e) \in \mathcal{E} \right\rbrace \right) \right)$$其中$\phi$ 是消息函数message function定义在每条边上将边特征与其两端节点特征组合生成消息$\rho$ 是归约函数reduce function将节点收到的所有入边消息聚合如sum、max、min、mean$\psi$ 是更新函数update function定义在每个节点上将聚合结果与节点自身特征结合并写回节点特征。这一范式贯穿 DGL 全部消息传递 API。原章正文见 message.rst其 Roadmap 将后续内容划分为四个主题内置函数与 API、高效编码、子图上的消息传递、异构图消息传递本文依次展开。内置函数与消息传递 APIEdgeBatch 与 NodeBatchUDF 的入参结构在 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是每个节点收到的消息数因此对消息求和时需对dim1归约。NodeBatch还提供nodes.data节点特征、nodes.nodes()节点 ID和nodes.batch_size()。更新函数同样接收nodes参数作用于归约函数的结果通常将其与节点原始特征组合并把结果保存为节点特征。优先使用 dgl.function 内置函数DGL 在命名空间dgl.function即fn下实现了常用消息函数与归约函数的内置版本built-in。DGL 官方建议只要可能就使用内置函数因为它们经过深度优化并且自动处理维度广播broadcasting。内置消息函数分为一元与二元两类一元unary支持copy例如copy_u、copy_e二元binary支持add、sub、mul、div、dot。命名约定为u代表源节点srcv代表目的节点dste代表边edge。参数均为字符串分别指定输入输出字段名。例如把源节点hu特征与目的节点hv特征相加、结果保存到边的he字段可写import dgl.function as fn fn.u_add_v(hu, hv, he)它等价于下面的消息 UDFdef message_func(edges): return {he: edges.src[hu] edges.dst[hv]}内置归约函数支持sum、max、min、mean源码见 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], dim1)}二元消息函数的完整集合在 python/dgl/function/message.py 中通过_register_builtin_message_func动态生成对目标组合u/v/e两两配对lhs ! rhs逐一注册add/sub/mul/div/dot五种运算因此实际可用函数包括u_add_v、u_mul_e、v_dot_e、e_sub_u等 30 个组合一元复制函数copy_u、copy_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消息传递一站式 APIupdate_all是高层次 API将消息生成、消息聚合、节点更新合并为一次调用从而为整体优化如内存复用留出空间。其签名为python/dgl/heterograph.pyupdate_all(message_func, reduce_func, apply_node_funcNone, etypeNone)三个核心参数分别为消息函数、归约函数与更新函数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})$$浮点类型支持与 float16DGL 内置函数支持浮点数据类型即特征必须是halffloat16/float/double张量。其中float16支持默认关闭因为它对 GPU 有最低算力要求计算能力需不低于sm_53即 Pascal、Volta、Turing 和 Ampere 架构。如需为混合精度训练启用 float16需要从源码编译 DGL具体步骤参见 Mixed Precision Training 教程。编写高效的消息传递代码DGL 对消息传递的内存消耗与计算速度做了专门优化。利用这些优化的常见做法是用内置函数作为参数将自定义消息传递逻辑组织成若干次update_all调用的组合。避免从节点到边的多余内存拷贝对于某些图边的数量远大于节点数量此时应尽量避免把节点特征拷贝到边上。但有些场景如dgl.nn.pytorch.conv.GATConvGAT 需要把消息保存在边上用于后续 softmax 等操作必须调用apply_edges配合内置函数在边上保存消息。由于边上的消息可能是高维的、非常耗内存DGL 建议尽可能保持边特征维度尽量低。下面是一个把边上的运算拆分到节点上执行的经典例子。目标是拼接源特征与目的特征再经过线性层即 $W \times (u \Vert v)$其中src、dst特征维度高而线性层输出维度低。直接实现低效——先拼接到边上再乘线性层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]], dim1)} 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_src和feat_dst保存在边上省内存且加法可以用 DGL 内置函数fn.u_add_v完成进一步加速计算、缩减内存占用。完整说明见 message-efficient.rst。在图的子图上应用消息传递如果只想更新图中的部分节点标准做法是先用节点 ID 构造子图再在子图上调用update_allnid [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。在图神经网络语境下根据复杂度不同某些节点类型和边类型可能需要用不同维数的表示来建模。异构图上消息传递可拆为两步对每个关系 r 分别做消息计算与聚合归约reduction把每个节点类型在所有关系上的聚合结果合并。DGL 在异构图上调用消息传递的接口是multi_update_allpython/dgl/heterograph.py。它接收两个参数字典以关系relation为键值为该关系下update_all的参数(message_func, reduce_func, [apply_node_func])字符串跨类型归约器cross type reducer可取sum、min、max、mean、stack。一个典型示例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异构图接口的完整体系。实践要点可归纳为优先内置函数性能优且自动广播UDF 仅在内置无法表达时使用更新函数外置将update_all外的更新写成纯张量运算代码更简洁降低边特征维度把拼接等重操作拆分到节点上执行减少节点→边的内存拷贝子图 update_all小批量训练的标准消息传递模式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),仅供参考
返回列表