PyTorch:非收缩维度下Tensordot的高效实现方法问询
高效实现张量按维度切片的类矩阵乘法操作
给定形状为(a, b, d)的张量x和形状为(b, c, d)的张量y,需要执行类矩阵乘法操作:收缩两者的第2维度(长度为d),但不对x的第1维度与y的第0维度(长度均为b)进行收缩,而是对该维度做切片迭代,最终得到形状为(a, b, c)的结果。
原实现采用Python循环逐次切片计算再堆叠,代码如下:
原循环实现代码:
import torch # 随机初始化参数 a = 3 b = 2 c = 5 d = 8 x = torch.randn(a, b, d) y = torch.randn(b, c, d) slice_results = [] for idx in range(b): x_slice = x[:, idx, :] y_slice = y[idx, :, :] slice_result = torch.tensordot(x_slice, y_slice, dims=([1], [1])) slice_results.append(slice_result) result = torch.stack(slice_results, dim=1) print(result.shape) # 输出: (3, 2, 5)
以下是两种无需显式构造列表的高效矢量化实现方式:
方法一:使用爱因斯坦求和约定(torch.einsum)
爱因斯坦求和可以直接用字符串描述张量的运算逻辑,代码简洁且直观,完全规避Python循环:
result = torch.einsum('abd,bcd->abc', x, y) print(result.shape) # 输出: (3, 2, 5)
逻辑说明:
'abd,bcd->abc'表示:取x的a/b/d维度、y的b/c/d维度- 对两者的
d维度执行求和收缩 - 最终保留
a/b/c维度,对应目标形状(a, b, c)
方法二:使用批量矩阵乘法(torch.bmm)
torch.bmm是专门针对三维张量的批量矩阵乘法,我们可以通过调整张量维度适配该API:
# 调整维度:将b作为batch维度,适配bmm的输入要求 x_reshaped = x.permute(1, 0, 2) # 形状从(a,b,d)转为(b,a,d) y_reshaped = y.permute(0, 2, 1) # 形状从(b,c,d)转为(b,d,c) # 执行批量矩阵乘法,再转置回目标形状 result = torch.bmm(x_reshaped, y_reshaped).permute(1, 0, 2) print(result.shape) # 输出: (3, 2, 5)
逻辑说明:
- 把
x的维度转置为(b,a,d),y转置为(b,d,c),此时b作为batch维度,每个batch对应原循环中的一组切片 bmm对每个batch执行矩阵乘法,得到形状为(b,a,c)的中间结果- 最后将中间结果转置为
(a,b,c),符合目标形状
以上两种方式均为PyTorch底层优化的矢量化操作,能充分利用CPU/GPU的并行计算能力,运算效率远高于原循环实现。
内容的提问来源于stack exchange,提问作者Sebastian Thomas
相关产品推荐
相关产品推荐

