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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 07:12:37