如何将PyTorch嵌套循环张量拼接改写为向量化实现
PyTorch 无循环向量化实现边特征拼接
可以直接利用PyTorch的张量维度扩展+广播机制,完全移除两层for循环,实现逻辑和原循环版本完全等价,运行效率远高于循环写法。
维度约定
- 输入张量:
edge_embeds:形状为(|A|, |B|, 2*edge_dim),存储A、B两类节点两两之间的原始边特征nodes_a_embeds:形状为(|A|, node_dim),存储A类所有节点的嵌入特征nodes_b_embeds:形状为(|B|, node_dim),存储B类所有节点的嵌入特征
- 目标输出:
edge_in:形状为(|A|, |B|, 2*node_dim + 2*edge_dim),每个位置对应A_i、B_u节点对的拼接特征,拼接顺序和原逻辑一致:A_i节点特征 + B_u节点特征 + 对应边原始特征
实现代码
# 无任何for循环的纯向量化实现 edge_in = torch.cat( [ nodes_a_embeds.unsqueeze(1), # 维度从(|A|, node_dim)扩展为(|A|, 1, node_dim),广播后自动对齐到(|A|, |B|, node_dim) nodes_b_embeds.unsqueeze(0), # 维度从(|B|, node_dim)扩展为(1, |B|, node_dim),广播后自动对齐到(|A|, |B|, node_dim) edge_embeds # 本身维度就是(|A|, |B|, 2*edge_dim),直接参与拼接 ], dim=-1 # 沿最后一个特征维度拼接 )
说明
- 不需要像原循环版本一样提前用
torch.ones初始化输出张量,torch.cat会直接申请对应大小的内存生成结果,减少冗余内存操作 - 该实现和原双层循环的计算结果完全一致,可以用小尺寸样例做逐元素对比验证
- 当|A|、|B|规模较大时,向量化实现的速度比循环版本高数十到数百倍,也避免了Python层循环的调度开销
内容的提问来源于stack exchange,提问作者emilector
相关产品推荐
相关产品推荐

