PyTorch实现无循环的多批次张量切片矩阵乘法
批量矩阵乘法实现(PyTorch)
现有两个PyTorch张量,形状定义如下:
A = [N x L x T] B = [N x T x K]
需求:对每个N维度对应的切片执行矩阵乘法(如A[0,:,:] @ B[0,:,:],得到形状[L x K]的结果),最终输出形状为[N, L, K]的张量,且禁止通过循环遍历N维度以保证计算效率。
解决方案
方法1:直接使用torch.matmul
PyTorch的torch.matmul原生支持批量矩阵乘法,会自动识别并保留前导的批量维度(此处为N),对每个批量内的子张量执行矩阵乘法:
result = torch.matmul(A, B) # 输出形状为[N, L, K],完全符合需求
方法2:使用torch.einsum(更直观的维度映射)
通过爱因斯坦求和符号明确指定维度对应关系,消除中间的T维度,直接得到目标形状:
result = torch.einsum('nlt,ntk->nlk', A, B)
- 符号说明:
nlt对应A的(N,L,T)维度,ntk对应B的(N,T,K)维度,箭头后的nlk表示输出维度为(N,L,K),其中T维度被求和抵消,与矩阵乘法逻辑一致。
验证示例
用小尺寸张量验证两种方法与循环结果的一致性:
import torch N, L, T, K = 2, 3, 4, 5 A = torch.randn(N, L, T) B = torch.randn(N, T, K) # 循环实现(用于对比) loop_result = torch.zeros(N, L, K) for i in range(N): loop_result[i] = A[i] @ B[i] # 方法1结果 matmul_result = torch.matmul(A, B) # 方法2结果 einsum_result = torch.einsum('nlt,ntk->nlk', A, B) # 检查一致性 print(torch.allclose(loop_result, matmul_result)) # 输出 True print(torch.allclose(loop_result, einsum_result)) # 输出 True
内容的提问来源于stack exchange,提问作者Inkyu Kim
相关产品推荐
相关产品推荐

