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

如何高效实现PyTorch张量与多个标量的逐元素乘法?

更优雅高效的PyTorch张量逐元素批量乘法实现

先看原始输入的张量:

x = torch.randint(1, 5, size=(2, 3, 3))
print(x.shape)
# torch.Size([2, 3, 3])

需要和以下标量张量做逐元素乘法,每个标量对应一次与x的乘法,最终把结果打包成一个张量:

weights = torch.tensor([2, 2, 2, 1])
print(weights.shape)
# torch.Size([4])

直接执行result = x * weights会因维度不匹配无法触发广播,而原方案通过repeat_interleave复制张量的方式既不优雅又低效:

x = x.unsqueeze(0).repeat_interleave(4, 0)
result = x * weights[:, None, None, None]

更优实现方法

利用PyTorch的广播机制,只需给weights添加合适的维度,无需复制x的内容,既节省内存又更简洁:

# 给weights添加3个维度,让它的形状变为(4,1,1,1),与x的(2,3,3)广播后匹配
result = x * weights.view(4, 1, 1, 1)
# 或者用unsqueeze链式调用,效果完全一致
result = x * weights.unsqueeze(1).unsqueeze(2).unsqueeze(3)

还可以用更简洁的索引写法扩展维度:

result = x * weights[..., None, None, None]

原理说明

核心是让weights的维度与x对齐:

  • weights原本是(4,),扩展后变为(4,1,1,1)
  • x是(2,3,3),广播时会自动将x做逻辑扩展(不会实际复制数据)为(4,2,3,3)
  • 逐元素乘法会在对应维度上完成,最终得到形状为(4,2,3,3)的结果张量,和原方案输出一致,但内存占用更低、效率更高。

内容的提问来源于stack exchange,提问作者kklaw

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 04:09:58