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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 18:45:05