基于PyTorch Geometric实现图节点特征的Softmax注意力池化
分图Softmax注意力池化实现(基于torch_scatter)
直接用torch_scatter的scatter_sum和scatter_max(可选,用于数值稳定)就能实现按batch分组的Softmax计算,具体步骤如下:
核心代码实现
import torch from torch_scatter import scatter_sum, scatter_max # 假设已有的输入: # scores: 形状(146,)的一维张量,每个节点的注意力分数 # batch: 形状(146,)的一维张量,标记每个节点所属的图索引(0到37) # --- 带数值稳定性的实现(推荐)--- # 1. 计算每个图内的最大分数,避免exp溢出 max_scores = scatter_max(scores, batch, dim=0)[0] # 形状(38,),每个图对应一个最大值 max_scores_per_node = max_scores[batch] # 扩展为(146,),每个节点对应所在图的最大值 # 2. 计算归一化后的exp值 exp_scores = torch.exp(scores - max_scores_per_node) # 3. 按图分组计算exp的和 sum_exp = scatter_sum(exp_scores, batch, dim=0) # 形状(38,),每个图的exp总和 sum_exp_per_node = sum_exp[batch] # 扩展为(146,),每个节点对应所在图的exp总和 # 4. 得到最终的注意力权重(每个图内的Softmax结果) attn_weights = exp_scores / sum_exp_per_node # --- 基础实现(无数值稳定,仅当scores范围较小时可用)--- # exp_scores = torch.exp(scores) # sum_exp = scatter_sum(exp_scores, batch, dim=0) # attn_weights = exp_scores / sum_exp[batch]
代码解释
scatter_max(scores, batch, dim=0):按batch分组,找出每个图内scores的最大值,用于后续数值稳定处理(防止指数运算导致的数值溢出)。max_scores[batch]:将每个图的最大值扩展到对应组内的所有节点位置,保证每个节点都减去所在图的最大值。scatter_sum(exp_scores, batch, dim=0):按batch分组,计算每个图内所有节点exp值的总和。sum_exp[batch]:将每个图的exp总和扩展到对应组内的所有节点位置,让每个节点的exp值除以所在图的总exp值,得到组内Softmax结果。
后续池化操作
得到注意力权重后,就可以用它对节点特征x(形状(146,256))进行池化:
pooled_features = scatter_sum(x * attn_weights.unsqueeze(-1), batch, dim=0) # pooled_features形状为(38,256),对应每个图的池化后特征
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

