PyTorch列维度成倍数关系的两个张量相加的更优实现方法问询
PyTorch张量批量相加高效实现方案
核心思路
利用PyTorch的广播机制和维度重排操作,完全替代Python层循环和张量拼接,所有运算都在张量底层执行。
实现代码
import torch # 示例张量定义 a = torch.tensor([[1., 1., 2.], [1., 1., 2.], [1., 1., 2.], [1., 1., 2.], [1., 1., 2.], [1., 1., 2.]]) b = torch.tensor([[4., 5., 6., 7., 8., 9.], [4., 5., 6., 7., 8., 9.], [4., 5., 6., 7., 8., 9.], [4., 5., 6., 7., 8., 9.], [4., 5., 6., 7., 8., 9.], [4., 5., 6., 7., 8., 9.]]) # 核心运算逻辑 k = b.shape[1] // a.shape[1] c = b.reshape(b.shape[0], k, a.shape[1]).add(a.unsqueeze(1)).reshape(b.shape)
方案优势
- 无Python层循环开销:所有运算都交由PyTorch底层实现,支持CPU/GPU加速,大张量场景下性能远高于循环拼接方案
- 内存效率更高:无需多次拼接中间张量,减少内存申请和拷贝开销
- 代码简洁易维护:仅需两行核心代码即可完成需求,适配任意满足
b.shape[1]是a.shape[1]整数倍的输入场景
结果验证
运行上述代码得到的c和预期输出完全一致:
tensor([[ 5., 6., 8., 8., 9., 11.], [ 5., 6., 8., 8., 9., 11.], [ 5., 6., 8., 8., 9., 11.], [ 5., 6., 8., 8., 9., 11.], [ 5., 6., 8., 8., 9., 11.], [ 5., 6., 8., 8., 9., 11.]])
内容的提问来源于stack exchange,提问作者hjsg1010
相关产品推荐
相关产品推荐

