异构图HAN

wen IT资讯 25

本文目录导读:

异构图HAN

  1. 背景:为什么要用 HAN?
  2. 核心架构:双层注意力机制
  3. 技术细节拆解
  4. 代码示例(简化版 DGL/PyG 风格)
  5. 学习建议

这是一个关于异构图注意力网络(HAN, Heterogeneous Graph Attention Network) 的深度解析,HAN 是处理异构图(包含多种类型节点和边的图)的一种经典且强大的深度学习模型。

背景:为什么要用 HAN?

在真实世界中,图数据很少是单一的。

  • 学术网络:包含作者、论文、会议、机构,节点类型不同,边的关系也不同(写作、引用、发表等)。
  • 电商网络:包含用户、商品、店铺、品牌。

传统图神经网络(如 GCN, GAT) 的局限性在于,它们默认所有节点和边是同质的,当把它们直接用在异构图上时,通常需要你手动把不同类型的信息“捏”成同样的特征空间(即投影),这忽略了异构性本身蕴含的丰富语义。

HAN 的核心洞察: 不同类型的关系(称为元路径)揭示了不同的语义信息,比如作者-论文-会议(A-P-C)这条路径表示“作者在某个会议发表论文”,而作者-论文-作者(A-P-A)表示“合作者关系”,HAN 的目的就是通过注意力机制,自动学习哪些元路径更重要,以及在一个元路径内部,哪些邻居节点更重要。

核心架构:双层注意力机制

HAN 包含两个层次的注意力机制:

  1. 节点级注意力:在一种元路径下,学习某个节点的邻居节点的重要性(谁更重要)。

    • 例子:通过“作者-论文-作者”这条元路径,学习该作者的哪些合作者对“他”当前的研究方向起决定性作用。
  2. 语义级注意力:在不同元路径之间,学习哪种元路径(语义)对于当前任务更重要。

    • 例子:在引文网络中,“合作者关系”(A-P-A)可能比“同机构关系”(A-Inst-A)更重要。

技术细节拆解

假设我们有节点类型 ( A )(作者)和 ( P )(论文),我们定义两条元路径:

  • ( \Phi_1 ):A-P-A(合作者)
  • ( \Phi_2 ):A-P-C-P-A(同主题作者)

第一步:节点级注意力

对于给定的元路径 ( \Phi ),所有节点都在该元路径的指导下形成一组节点对(即邻居),我们使用一个自注意力机制来计算节点 ( j ) 对节点 ( i ) 的重要性。

  1. 转换:由于不同类型的节点特征维度可能不同,首先对输入特征 ( h_i ) 进行线性变换(通过一个类型特定的映射矩阵),投影到统一的隐含空间。
  2. 注意力计算:计算节点 ( j ) 对 ( i ) 的注意力系数: [ e{ij}^{\Phi} = att{\text{node}}(h_i', hj'; \Phi) ] 这里 ( att{\text{node}} ) 通常是一个单层的前馈网络(类似 GAT)。
  3. Softmax 归一化:在节点 ( i ) 的所有邻居 ( \mathcal{N}i^{\Phi} ) 上进行归一化: [ \alpha{ij}^{\Phi} = \frac{\exp(\sigma(\mathbf{a}_{\Phi}^T \cdot [h_i' \,||\, hj']))}{\sum{k \in \mathcal{N}i^{\Phi}} \exp(\sigma(\mathbf{a}{\Phi}^T \cdot [h_i' \,||\, h_k']))} ]
  4. 聚合:通过加权求和得到节点 ( i ) 在元路径 ( \Phi ) 下的语义嵌入 ( z_i^{\Phi} )。

第二步:语义级注意力

现在我们有了多个元路径下的节点表示(( Z_{\Phi1}, Z{\Phi_2} )),需要将它们融合。

  1. 重要性计算:假设所有节点共享同一个元路径注意力向量 ( \mathbf{q} ),计算每个元路径 ( \Phii ) 的重要性: [ w{\Phii} = \frac{1}{|V|} \sum{i \in V} \mathbf{q}^T \cdot \tanh(\mathbf{W} \cdot z_i^{\Phi_i} + \mathbf{b}) ] 这里对全图所有节点的嵌入取平均(全局上下文)来计算该元路径的权重。
  2. Softmax 归一化: [ \beta_{\Phii} = \frac{\exp(w{\Phii})}{\sum{k=1}^{M} \exp(w_{\Phi_k})} ]
  3. 最终嵌入: [ Z = \sum{k=1}^{M} \beta{\Phik} \cdot Z{\Phi_k} ]

最终的 ( Z ) 就是融合了多种语义信息的节点嵌入,可用于下游任务(节点分类、聚类、链接预测等)。

代码示例(简化版 DGL/PyG 风格)

由于 HAN 的实现涉及自定义元路径,这里提供一个使用 DGL 库构建 HAN 的核心逻辑示意(非完整运行代码,重点在逻辑):

import 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):
        super().__init__()
        self.attn_fc = nn.Linear(2 * hidden_dim, 1, bias=False)
    def forward(self, g, h):
        # g 是元路径对应的同构图(邻接矩阵)
        with g.local_scope():
            g.ndata['h'] = h
            # 通过边缘计算注意力分数
            g.apply_edges(lambda edges: {
                'e': self.attn_fc(torch.cat([edges.src['h'], edges.dst['h']], dim=1))
            })
            # Softmax 归一化
            g.edata['a'] = dgl.softmax_nodes(g, 'e')
            g.ndata['z'] = dgl.sum_edges(g, 'a' * 'h')
            return g.ndata['z']
class SemanticLevelAttention(nn.Module):
    """语义级注意力融合多条元路径"""
    def __init__(self, in_dim, num_metapaths):
        super().__init__()
        self.fc = nn.Linear(in_dim, in_dim)
        self.query = nn.Parameter(torch.randn(1, in_dim))
    def forward(self, h_list):
        # h_list: 包含每条元路径下的节点嵌入列表
        # 计算每条元路径的全局注意力权重
        w = []
        for h in h_list:
            h_mean = h.mean(dim=0)  # 全图平均
            w_i = (self.query @ torch.tanh(self.fc(h_mean)).T).squeeze()
            w.append(w_i)
        beta = F.softmax(torch.stack(w), dim=0)
        # 加权融合
        z = sum(beta[i] * h_list[i] for i in range(len(h_list)))
        return z, beta
class HAN(nn.Module):
    def __init__(self, meta_paths, in_dim, hidden_dim, out_dim):
        super().__init__()
        self.meta_paths = meta_paths
        # 每个元路径对应一个节点级注意力
        self.node_attns = nn.ModuleList([
            NodeLevelAttention(in_dim, hidden_dim) for _ in meta_paths
        ])
        self.semantic_attn = SemanticLevelAttention(hidden_dim, len(meta_paths))
        self.fc = nn.Linear(hidden_dim, out_dim)
    def forward(self, g_list, h_dict):
        # g_list: 每条元路径转换后的同构图列表
        # h_dict: 原始异构图中各节点的初始特征
        meta_embeds = []
        for i, g in enumerate(g_list):
            # 这里假设所有类型节点已经投影到统一空间
            h = ...  # 通过类型转换获取特征
            z = self.node_attns[i](g, h)
            meta_embeds.append(z)
        # 语义融合
        z_final, beta = self.semantic_attn(meta_embeds)
        return self.fc(z_final)

学习建议

  1. 先理解元路径:这是 HAN 的硬性前提,你需要根据领域知识定义有意义的元路径。

    • 社交网络:U-U(好友), U-G-U(同群组)。
    • 推荐系统:U-I-U(用户看过同类商品), U-S-I(用户关注店铺的商品)。
  2. 对比 GAT:HAN 的节点级注意力类似于 GAT,但针对不同的元路径分别做,而且增加了语义融合层。

  3. 挑战

    • 元路径的选择:这是个艺术活,选不好效果会差。
    • 可扩展性:全图计算语义级权重(对所有节点嵌入取平均)在大图上可能受限,后续改进版本(如 HAN-Sampling)通过子图采样来解决。
  4. 训练技巧:如果你的类别不平衡,建议采用 元路径采样(Metapath2vec 提出)来生成负样本,而不是简单的随机采样。

总结一句话:HAN 通过“节点级注意力”在每条语义路径内部选边,再通过“语义级注意力”在不同语义路径之间择优,从而动态捕捉异构图中的复杂交互关系。

上一篇知识图嵌入

下一篇图对比学习

抱歉,评论功能暂时关闭!