基于PyTorch的无循环散点矩阵乘法GPU并行化实现需求
向量化实现分组矩阵转置乘法(GPU友好)
原代码通过循环对每个graph分组的节点特征计算pred_m.T @ pred_m,在GPU环境下循环会带来不必要的调度开销。我们可以利用向量化操作+分组求和的方式完全消除循环,充分利用GPU的并行计算能力:
import torch node_predictions = torch.randn(size=(15, 8), device="cuda") # 直接部署到GPU node2graph = torch.tensor([0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 2], dtype=torch.int64, device="cuda") # 步骤1:计算所有节点特征的外积(对应单个节点对目标矩阵的贡献) # 输出shape: (num_nodes, feat_dim, feat_dim) = (15,8,8) node_outer = node_predictions.unsqueeze(2) * node_predictions.unsqueeze(1) # 步骤2:按graph分组求和,得到每个graph的 pred_m.T @ pred_m num_graphs = node2graph.max().item() + 1 # 初始化结果张量,与原代码输出shape一致:(num_graphs,8,8) pred_m_T_times_pred_m = torch.zeros(num_graphs, 8, 8, device=node_predictions.device) # 按node2graph的索引分组累加外积 pred_m_T_times_pred_m.scatter_add_( dim=0, index=node2graph.unsqueeze(1).unsqueeze(2).expand(-1, 8, 8), # 扩展索引以匹配外积张量shape src=node_outer ) # 验证与原代码结果一致性(仅用于校验) split_size = torch.bincount(node2graph) list_of_node_predictions = torch.split(node_predictions, split_size.tolist()) original_result = torch.stack([torch.matmul(pred_m.T, pred_m) for pred_m in list_of_node_predictions]) print(torch.allclose(pred_m_T_times_pred_m, original_result)) # 输出True
核心原理
pred_m.T @ pred_m本质是对分组内所有节点特征向量的外积求和:
$$\sum_{i \in \text{graph}_k} v_i v_i^T$$
我们先一次性计算所有节点的外积,再通过scatter_add_按graph索引批量累加,完全避免了循环,GPU可以并行处理所有外积计算和分组求和操作,效率远高于循环实现。
优势
- 无循环,完全释放GPU并行计算能力
- 内存使用高效,无需拆分张量生成中间列表
- 代码简洁,适配大规模graph数据场景
内容的提问来源于stack exchange,提问作者Trf
相关产品推荐
相关产品推荐

