ARTICLE DETAIL

资讯详情

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

HAN异构图注意力网络:双层注意力机制与元路径设计实战

HAN异构图注意力网络:双层注意力机制与元路径设计实战 图学习系列写到第四篇终于轮到HAN异构图注意力网络。如果你已经熟悉GCN、GAT第一次看到HAN时大概率会觉得这不就是在GAT外面再套一层注意力吗我最初也有这种错觉直到在项目里真正处理学术网络数据时才发现这“多出来的一层”解决的是一个结构性难题——节点和边都有类型类型不同语义就不同一套共享参数的GAT根本扛不住。这篇笔记围绕HAN的动机、两层注意力机制、工程实现细节和实际踩坑展开适合已经跑通过GCN/GAT、现在准备把图模型迁移到异构图数据上的读者。1. 为什么同构图注意力解决不了异构图问题一个学术网络的例子1.1 我手上的业务问题长什么样先说一个很常规的场景给定一个学术网络里边有作者A、论文P、期刊/会议V、主题词T四类节点。任务是对作者的研究方向做分类。这时候你把所有节点和边简单塞给GCN得到的结果通常不好不是因为GCN本身弱而是因为这种图上的信息天然是“分类型”的作者和作者之间没有直接边他们通过合著论文产生联系作者和期刊/会议之间隔着论文但“在哪个会议发文”对研究方向判断非常关键作者和主题词之间也隔着论文主题词又能反映研究内容的语义相似性。如果你用GCN它默认所有邻居一视同仁作者节点会把“合著的作者”“发表论文的会议”“论文主题词”全部揉在一起更新。可问题是这三类邻居对“作者研究方向”的贡献方式完全不同。GCN的传播规则里没有“关系类型”这个概念你把异构关系强行转成同构的一阶邻接矩阵等于在输入端把关键信息丢了一大半。GAT稍微好一点它会为每个邻居学一个注意力权重但它学的是“哪个邻居重要”不是“哪一类关系应该用什么样的方式聚合”。换个说法GAT可以在聚合时给某个会议节点更高权重但它无法区分“这个节点是通过作者-论文-会议链路到达的”和“这个节点是作者直接关注的某个人”在语义上的差异。1.2 当“异质”叠加“注意力”难点到底在哪异构图的“异”体现在三个层面。第一节点类型集合大于1。学术网络里至少有作者、论文、会议、主题词四种节点。不同类别节点的特征空间还不一样作者可以是ID embedding论文是一段文本编码会议可以用向量表示主题词又是另一套分布。你没法简单地把它们放进同一个特征矩阵。第二关系类型大于1。作者到论文是“写作”论文到会议是“发表”论文到主题词是“包含”。这些关系不是简单的高低权重而是语义类型上的区别。“作者A写论文P”和“论文P收录于会议V”是两种完全不同的信息通道。第三路径和路径之间还会产生复合语义。作者A和作者B可以通过“A写论文PB也写论文P”建立联系合著关系也可以通过“A写论文PP发表在会议VB也有论文发表在V”建立联系同会议关系。这两种路径对“A和B是否属于同一研究方向”的判断贡献完全不同。同构图模型处理不了这些因为它在结构上就不支持“类型化路径”。HAN做了一件很聪明的事先把复杂的异构图按“元路径”拆成多个同构子图在每个子图内部用节点级注意力做邻居聚合再用一个语义级注意力把多个子图的结果融合起来。这样既保留了类型信息又没有增加暴力建模的复杂度。2. HAN的工作方式拆解两层注意力各管一段语义2.1 节点级注意力在一条元路径内部筛选重要邻居理解HAN的第一步是先理解元路径。元路径是连接两个同类型节点的复合路径它定义了一种“视角”。以作者分类任务为例常用的元路径有A-P-A两个作者合作过同一篇论文合著视角A-P-V-P-A两个作者在同一个会议/期刊发表过论文同领域视角A-P-T-P-A两个作者的研究主题词有重叠主题相似视角。每一条元路径都定义了图中的一种同构关系。对作者节点A1来说在“A-P-A”这条路径下它的“可达邻居”是那些和他合著过论文的作者在“A-P-V-P-A”下它的邻居变成了在同一会议发过文的作者。HAN的节点级注意力要做的事情就是针对每一条元路径在对应的邻居集合上计算注意力权重然后加权聚合该路径下的节点表示。写成公式形式设节点 i 在第 Φ 条元路径下的邻居集合为 (N_i^{\Phi})先把节点原始特征通过一个类型相关的线性变换投影到统一维度 (h_i W_{\Phi} x_i)然后计算注意力系数[ \alpha_{ij}^{\Phi} \frac{\exp\left(\sigma\left(a_{\Phi}^T [h_i | h_j]\right)\right)}{\sum_{k \in N_i^{\Phi}} \exp\left(\sigma\left(a_{\Phi}^T [h_i | h_k]\right)\right)} ]其中 (\sigma) 是 LeakyReLU 激活(|) 是拼接操作。这个系数表示在这条元路径视角下邻居 j 对节点 i 有多重要。然后用这个系数加权求和得到节点 i 在路径 Φ 下的表示[ z_i^{\Phi} \sum_{j \in N_i^{\Phi}} \alpha_{ij}^{\Phi} h_j ]这里的核心是“每条元路径都单独维护一套注意力参数 (a_{\Phi}) 和投影矩阵 (W_{\Phi})”。也就是说不同视角下的“重要邻居”是分开学的会议视角的注意力不会污染合著视角的注意力。2.2 语义级注意力在多个元路径之间做权衡一条元路径只能表达一种语义。实际任务里我们通常同时选多条元路径比如上面那三条。这时候每个节点会得到三组表示(Z_{A-P-A})、(Z_{A-P-V-P-A})、(Z_{A-P-T-P-A})。问题来了最终做作者分类时这三位“专家”的意见谁更可信对不同作者来说三条路径的权重可能不同——有些人靠合著关系就足够区分研究方向有些人需要看会议类型还有些人更强依赖于主题词相似。HAN的第二层注意力就是来解决这个问题的。它把这多条元路径看作多个“语义视图”学习每个视图的全局重要性权重[ \beta_{\Phi} \frac{\exp\left(w_{\text{sem}}^T \tanh(W_{\text{sem}} Z^{\Phi})\right)}{\sum_{\Phi} \exp\left(w_{\text{sem}}^T \tanh(W_{\text{sem}} Z^{\Phi})\right)} ]这里的 (w_{\text{sem}}) 和 (W_{\text{sem}}) 是语义级注意力的可学习参数(Z^{\Phi}) 是某条元路径下所有节点表示组成的矩阵。把每条路径的全局权重和对应的节点表示加权求和就得到最终的节点表示[ Z \sum_{\Phi} \beta_{\Phi} Z^{\Phi} ]注意语义级注意力学到的是“全局权重”不是逐节点的。也就是说模型假设对整批训练数据来说某条元路径对任务的整体贡献是相对稳定的。这个设计在学术网络的作者分类任务上合理因为作者的研究方向分布会使得某些关系模式整体上更有判别力。2.3 为什么要“先分后合”而不是设计一个全局注意力有个很自然的疑问为什么不能直接把异构邻接矩阵丢给一个多关系感知的注意力模型非要拆成元路径再融合我自己的理解是这背后有三个考虑。一是复杂度问题。全局注意力要做两两节点之间的注意力计算复杂度是 (O(N^2))。学术网络动辄几十万节点直接全图注意力根本不现实。HAN通过元路径采样把注意力计算限制在了“该路径可达的邻居”内复杂度降为 (O(\sum |E_{\Phi}|))和同构图GAT相当。二是参数量的考虑。异构图里关系类型可能非常多如果为每种边单独建模一套传播参数RGCN就是这么干的参数会随关系数量线性膨胀在小数据集上很容易过拟合。HAN把无限多种关系压缩到有限的几条元路径里每条路径共享一套参数参数量可控得多。三是可解释性。语义级注意力输出的是每条元路径的权重你可以直接看到模型在做分类时更依赖合著关系还是主题相似关系。这个输出对业务分析很有用纯GAT给不出这种解释。3. 设计决定效果元路径、特征投影、多头与训练细节3.1 元路径选型从任务倒推候选路径很多新手把HAN跑不好的第一个原因是元路径选得太随意。这里我给几条经验。先看任务要预测什么。节点分类任务里元路径的两端通常是你希望建模的目标节点类型中间节点类型决定语义。比如IMDB电影分类任务目标节点是电影M候选元路径M-A-M表示“两个电影有共同演员”M-D-M表示“两个电影有共同导演”。演员和导演对电影类别的判别力不同所以这两条路径都有保留价值。路径不宜太长。元路径长度超过4跳后噪声通常会大于信号。我做实验时发现A-P-T-P-A4跳比A-P-A2跳在部分指标上有提升但A-P-V-P-T-P-V-P-A这种7、8跳的路径几乎不涨点反而拖慢训练。路径越长经过中间节点聚合后端节点之间的相关性越稀薄这是一种“稀释效应”。路径数量控制在2到4条。有人觉得路径越多信息越全实际不是。语义级注意力本质上是在和高层语义交互路径太多会稀释每条路径的监督信号注意力权重也很难收敛。我在DBLP子集上的经验是3条路径就能达到不错的分类效果再加路径收益很小偶尔还掉点。3.2 类型投影和维度对齐最容易忽略的细节HAN要求不同元路径输出的节点表示维度一致否则没法做加权求和。这个“维度一致”通常靠类型相关的投影矩阵 (W_{\Phi}) 保证。具体做法是每条元路径关联一个线性层把该路径下涉及的节点特征投影到同一个隐藏维度。比如作者节点的原始特征是词袋向量会议节点可能是one-hot或者embedding二者维度差异很大。投影层的作用不只是“对齐维度”更重要的是在每个元路径语义空间里重新刻画节点。我建议不同元路径使用不同的投影矩阵而不是共享同一个。原因在于A-P-A的投影空间服务于“作者-作者协作关系”A-P-V-P-A的投影空间服务于“会议发表关系”两者强调的特征维度不同。共享投影会把这两类语义强行挤压进同一个线性空间模型表达能力受限。维度选择上隐藏层维度64、128是比较稳的起点。维度过低比如16注意力打分能表达的空间太小分类指标明显下滑维度过高比如512在论文规模的数据集上容易过拟合。3.3 多头注意力与关键超参HAN在节点级注意力部分沿用了多头设计。每个头独立计算注意力系数输出可以拼接或取平均。我的使用习惯是中间层拼接最后一层取平均。拼接能保留每个头的特征子空间有利于信息表达最后一层输出要进入分类器取平均可以抑制不同头之间的方差让分类层更稳定。这两者的差异在小数据集上比较明显。超参方面直接给一组经过验证的基准配置隐藏维度64注意力头数8dropout 0.5学习率0.005weight decay 0.001优化器Adam训练epoch 100到200之间并配合早停。这组参数在DBLP、ACM、IMDB三个经典异构图数据集上都能稳定工作。如果你的图特别大或者特征维度特别高可以适当调低学习率到0.001并把dropout提高到0.6。4. 最小可复现的HAN训练流程PyTorch核心代码4.1 数据组织邻接表与元路径矩阵HAN的训练数据组织起来比GAT略复杂因为每条元路径都需要一个单独的邻接关系。我一般先把图存成邻接表然后用预计算好的元路径实例生成多个稀疏邻接矩阵。以A-P-A为例先找到所有“作者-论文”对再按共同论文对作者进行配对。这一步可以直接用矩阵乘法实现假设矩阵 (M_{AP}) 表示作者-论文关系那么作者之间的合著矩阵就是 (M_{AP} M_{AP}^T)。A-P-V-P-A对应的邻接矩阵则是 (M_{AP} M_{PV} M_{VP} M_{PA})。实际操作中如果图很大这步可能要分块计算避免一口气把稠密矩阵存进内存。生成好每条元路径的邻接矩阵后后续训练就是一个“多通道同构图”任务每个通道对应一条元路径通道内部跑GAT式聚合通道之间再套一层注意力。4.2 节点级模块和语义级模块的代码骨架这里给一份按核心逻辑精简过的PyTorch实现能帮你把HAN的结构落到代码层面。正式项目里可以把稀疏矩阵运算换成PyG或DGL的内置算子效率更高。import torch import torch.nn as nn import torch.nn.functional as F class NodeLevelAttention(nn.Module): def __init__(self, in_dim, hidden_dim, n_heads, dropout0.5): super().__init__() self.hidden_dim hidden_dim self.n_heads n_heads self.W nn.Linear(in_dim, hidden_dim * n_heads, biasFalse) self.a nn.Parameter(torch.empty(n_heads, 2 * hidden_dim)) nn.init.xavier_uniform_(self.a) self.leaky_relu nn.LeakyReLU(0.2) self.dropout nn.Dropout(dropout) def forward(self, x, adj): # x: [N, in_dim]adj: [N, N] 稀疏邻接矩阵 N x.size(0) h self.W(x).view(N, self.n_heads, self.hidden_dim) edge_index adj._indices() # [2, E] src, dst edge_index[0], edge_index[1] outputs [] for head in range(self.n_heads): h_head h[:, head, :] # [N, hidden_dim] h_src h_head[src] h_dst h_head[dst] beta torch.cat([h_src, h_dst], dim-1) score self.leaky_relu((beta * self.a[head]).sum(dim-1)) score torch.exp(score) # 归一化按目标节点聚合 msg torch.zeros(N, devicex.device) msg.index_add_(0, dst, score) score_norm score / (msg[dst] 1e-8) h_agg torch.zeros(N, self.hidden_dim, devicex.device) h_agg.index_add_(0, dst, h_src * score_norm.unsqueeze(-1)) outputs.append(h_agg) return torch.cat(outputs, dim-1) # [N, hidden_dim * n_heads] class SemanticAttention(nn.Module): def __init__(self, in_dim, hidden_dim128): super().__init__() self.fc1 nn.Linear(in_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, 1) def forward(self, z_list): # z_list: 每条元路径下的节点表示每个是 [N, in_dim] scores [] for z in z_list: s self.fc2(torch.tanh(self.fc1(z))) # [N, 1] scores.append(s) weight torch.softmax(torch.stack(scores, dim0), dim0) # [P, N, 1] zs torch.stack(z_list, dim0) # [P, N, in_dim] return (weight * zs).sum(dim0)上面的index_add_方式不是最高效的但对理解HAN的语义足够清晰。实际训练时我建议把注意力计算放到稠密block上或者用PyG的MessagePassing重写速度能快一个量级。4.3 训练循环与评估指标HAN训练循环和普通GNN分类任务没有本质区别关键是每一轮都要循环所有元路径分别做节点级聚合再做语义级融合最后过分类器算交叉熵。一个简化的训练步骤长这样def train_step(model, x_dict, adj_dict, y, mask): model.train() optimizer.zero_grad() z_list [] for path_name, adj in adj_dict.items(): x x_dict[path_name] # 该路径下的节点特征 z model.node_level(x, adj) # 节点级注意力 z_list.append(z) z_final model.semantic_level(z_list) # 语义级注意力 logits model.classifier(z_final) loss F.cross_entropy(logits[mask], y[mask]) loss.backward() optimizer.step() return loss.item()评估指标我习惯同时报macro-F1和micro-F1。学术网络分类数据集的类别通常不平衡比如ACM数据里某些论文类别样本明显偏少macro-F1能更真实地反映少数类效果。如果你只盯着accuracy很容易把模型优化成对多数类过拟合而忽略少数类。5. 我在实际项目里踩过坑之后总结的几条经验5.1 数据划分的坑掩码别乱用类别要均衡HAN原始论文是transductive设置训练、验证、测试节点是从同一个图里划分出来的。这带来一个容易被忽略的问题验证集和测试集的节点也参与训练时的邻居聚合只是标签被掩码掉。这不是bug是设计如此但如果你不太熟悉这一点会让复现实验的结果和直觉对不上。我在复现ACM数据集时第一次跑出来的macro-F1和论文对不上排查发现是数据划分不均匀某些类别在训练集里只有一个节点模型几乎学不到该类别的特征。后来复现代码时采用了随机分层抽样保证每个类别在训练验证测试三部分中按比例分配指标立刻正常了。另外一个容易踩的坑是节点ID顺序做mask时如果不和数据集的原生ID对齐会出现“训练标签错位”的问题表现为主loss能降但验证集乱跳。5.2 过平滑与“注意力退化”问题很多人以为HAN这么复杂的结构总该能叠很多层吧实际上不是。HAN在节点级聚合部分本质上还是在做“邻居表示平均”的变体层数一旦超过3层同样的过平滑问题就会出现所有节点的表示趋向一致分类边界消失。我的判断方法是每训练完一轮抽样几百个节点计算相邻两层的输出向量平均余弦相似度。如果相似度超过0.95基本可以判定过平滑已经开始影响效果这时候应该减层而不是加层。还有一个现象叫“注意力退化”我看过一些文章没提。语义级注意力在训练后期可能逐渐收敛成近似one-hot分布比如只信任A-P-A路径其他路径权重趋近于0。表面上看,这是模型在做特征选择实际上是因为某一类元路径含有的监督信号太强把其他路径的梯度压没了。缓解办法有两个一是给语义级注意力权重加一个熵正则项鼓励它在早期保持分散二是每条元路径单独接一个分类辅助任务给弱路径额外的梯度来源。5.3 冷启动、批次采样和内存问题HAN原始训练方式是全图transductive这在初始的DBLP规模上没问题但一旦换成千万节点级别的工业图全图计算就吃不消了。节点级注意力的每一条元路径都要维护一个注意力系数矩阵内存开销会随边数线性增长。我试过一次几百万条边的元路径单卡A100直接OOM。工业场景里更可行的做法是引入mini-batch训练就是一个batch里采样一部分节点然后为这些节点的元路径邻居做固定跳数的采样。这样一来模型变成inductive还能顺带解决新节点冷启动问题——新节点只要有边关系就能通过采样聚合出表示不需要重新训练。但要注意换了采样策略后HAN里的语义级注意力会变“不稳定”因为每个batch里不同元路径的邻居数量可能差异很大权重容易偏向“采样到更多邻居”的路径。解决办法是做归一化处理或者对每条路径的注意力输入做LayerNorm我在项目里用后者效果更稳定。6. HAN和周边模型的关系以及我现在的选型建议6.1 同GCN/GAT/RGCN/HGT的横向比较我整理了一张表方便对照选择模型模型能否处理节点类型能否处理边类型语义可解释性参数量典型应用场景GCN否否无低同构图节点分类GAT否否节点重要性中同构图、注意力需求RGCN否是按关系区分随关系数膨胀知识图谱补全HAN是元路径内是元路径级元路径权重清晰中异构图节点分类HGT是是较弱高大规模异构图、自动关系建模HAN的优势在于“中等复杂度 强可解释性”。它不需要像RGCN那样为每种关系单独建参也不会像HGT那样依赖复杂的位置编码和动态注意力。代价是你必须人工设计元路径这既是限制也是可控性来源。6.2 什么场景值得优先考虑HAN我的选型建议分三种情况。如果你的图确实是同构的或者异构关系非常稀疏直接用GAT或GraphSAGE就好HAN带来的提升有限还会引入元路径设计成本。如果图的关系类型在5到10种之间且你能根据领域知识给出2到4条有意义的元路径HAN是目前性价比最高的选择。它在DBLP、ACM、IMDB这些标准数据集上反复验证过训练稳定调参空间小。如果关系类型非常多且难以人工归纳元路径我会改用HGT或基于Transformer的异构图模型。它们能从数据里自动学习类型间交互但对数据量和算力的要求明显更高。6.3 从HAN延伸出去的方向图强化学习HAN学到的高质量节点表示不只能用来做分类。我在后续项目里把它当作编码器接到深度强化学习框架里让智能体在异构图环境中基于HAN生成的节点状态做决策。比如在推荐场景里用户和物品是两类节点交互行为是异构边HAN可以先把用户状态编码成向量再交给策略网络决定推荐动作整套流程比直接用同构图编码器更符合业务直觉。图强化学习和深度强化学习结合是目前比较活跃的方向HAN在其中承担的“表示学习”角色很清晰如果你对决策方向感兴趣可以从这个角度切入。最后再分享一个我个人的实操体会HAN是一种“设计大于调参”的模型。很多人花了大量时间调学习率、dropout效果上不去回头检查才发现是元路径设计不符合任务语义。先花足够的精力把图的结构、任务的目标、候选元路径梳理清楚再用HAN往往比你盲目堆模型层数有效得多。把这套思路跑通一次再看其他异构图模型心里会踏实很多。
返回列表