PyTorch嵌套循环性能优化求助:尝试广播优化未果
优化PyTorch嵌套循环的方案
你的核心问题是用Python嵌套循环处理批量样本和图结构计算,这会带来极大的性能开销。我们可以通过矩阵乘法替代循环,利用PyTorch的批量运算能力彻底消除所有嵌套循环,同时处理mask的分支逻辑。
先拆解原代码的计算逻辑
原循环里的计算可以转化为线性代数运算:
- 当
mask[m]==1时,student_output[m][i] = sum_j (student_adjacency_mat[i][j] * dot(student_feat[m], student_graph[j]))
等价于:student_output[m] = student_feat[m] @ (student_adjacency_mat @ student_graph).T - 当
mask[m]==0时,teacher_output[m][i] = sum_j (teacher_adjacency_mat[i][j] * dot(teacher_feat[m], teacher_graph[j]))
等价于:teacher_output[m] = teacher_feat[m] @ (teacher_adjacency_mat @ teacher_graph).T
优化后的完整代码
import torch batch_size=32 mask=torch.FloatTensor(batch_size).uniform_() > 0.8 mask_bool = mask.bool() # 转成布尔型方便索引 teacher_count=510 student_count=420 feature_dim=750 # 初始化输出 student_output=torch.zeros([batch_size,student_count]) teacher_output=torch.zeros([batch_size,teacher_count]) student_adjacency_mat=torch.randint(0,1,(student_count,student_count)) teacher_adjacency_mat=torch.randint(0,1,(teacher_count,teacher_count)) student_feat=torch.rand([batch_size,feature_dim]) student_graph=torch.rand([student_count,feature_dim]) teacher_feat=torch.rand([batch_size,feature_dim]) teacher_graph=torch.rand([teacher_count,feature_dim]) # 1. 预计算邻接矩阵与图特征的乘积,消除j维度的循环 student_adj_graph = student_adjacency_mat @ student_graph # shape: (student_count, feature_dim) teacher_adj_graph = teacher_adjacency_mat @ teacher_graph # shape: (teacher_count, feature_dim) # 2. 批量计算所有样本的输出,消除m和i维度的循环 student_output_batch = student_feat @ student_adj_graph.T # shape: (batch_size, student_count) teacher_output_batch = teacher_feat @ teacher_adj_graph.T # shape: (batch_size, teacher_count) # 3. 根据mask赋值到最终输出 student_output[mask_bool] = student_output_batch[mask_bool] teacher_output[~mask_bool] = teacher_output_batch[~mask_bool]
优化效果说明
- 完全消除了三重嵌套循环,所有计算都用PyTorch的底层优化算子完成(CPU下用BLAS,GPU下用CUDA核),速度会提升几个数量级。
- 批量处理所有样本,不需要逐样本判断mask,利用布尔索引直接完成分支赋值。
- 计算逻辑和原代码完全一致,没有改变结果。
内容的提问来源于stack exchange,提问作者Aleph
相关产品推荐
相关产品推荐

