如何高效实现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
相关产品推荐
相关产品推荐

