You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.10 12:05:09