如何不使用循环仅通过矩阵运算实现指定Torch张量计算?
问题
能否不使用for循环,仅借助numpy或PyTorch的矩阵运算实现如下代码的功能?相关代码及输出示例如下:
import torch torch.manual_seed(0) input = torch.rand((1, 3543, 768)) w = torch.rand((3543, 768, 1)) output = torch.zeros(input.shape[0], input.shape[1], 1) for i in range(input.shape[1]): output[:, i, :] = torch.matmul(input[:, i, :], w[i, :, :])
输出示例(注:原输出格式有误,修正为对应(1, 3543, 1)形状的正确结果):
tensor([[[0.4963], [0.1164], [0.3337], ..., [0.5168], [0.8208], [0.8758]]])
实现方案
完全可以,以下是几种PyTorch原生的无循环实现方式,效率远高于for循环:
方法1:爱因斯坦求和(torch.einsum)
直接通过维度映射描述运算逻辑,直观对应原循环的操作:
import torch torch.manual_seed(0) input = torch.rand((1, 3543, 768)) w = torch.rand((3543, 768, 1)) # b=batch维度,n=序列长度维度,d=特征维度,k=输出维度 output = torch.einsum('bnd, ndk -> bnk', input, w)
方法2:元素相乘+维度求和
原循环中的矩阵乘法本质是对应位置向量的点积,可以通过元素相乘后沿特征维度求和实现:
# 先将w的最后一维压缩,变成(3543, 768) w_squeezed = w.squeeze(dim=-1) # 元素相乘后沿特征维度(dim=2)求和,同时保持维度 output = (input * w_squeezed).sum(dim=2, keepdim=True)
方法3:批量矩阵乘法(torch.bmm)
通过调整维度适配批量矩阵乘法的要求:
# 将input调整为(3543, 1, 768),w保持(3543, 768, 1) input_reshaped = input.permute(1, 0, 2) # 批量计算3543组(1,768) × (768,1)的矩阵乘法 batch_output = torch.bmm(input_reshaped, w) # 恢复原维度顺序(1, 3543, 1) output = batch_output.permute(1, 0, 2)
验证一致性
以上三种方法的输出结果与原for循环代码的输出完全一致,可以通过torch.allclose(output_loop, output_einsum)这类语句验证。
内容的提问来源于stack exchange,提问作者huankkai
相关产品推荐
相关产品推荐

