如何用einsum实现批量向量的x@A^T@A@x^T@x@A@A^T@x^T计算
问题:矩阵运算的Einstein求和(einsum)实现
现有m个形状为(1,n)的向量x,横向堆叠为形状(m,n)的矩阵B;另有一个(n,n)的矩阵A。需要为每个x计算表达式 x@A^T@A@x^T@x@A@A^T@x^T,最终输出形状为(m,1)的结果,求对应的einsum查询语句。
用户提供的非einsum实现示例代码:
import torch m = 30 n = 4 B = torch.randn(m, n) A = torch.randn(n, n) result = torch.zeros(m,1) for i in range(m): x = B[i].unsqueeze(0) result[i] = torch.matmul(x,torch.matmul(A.T,torch.matmul(A,torch.matmul(x.T,torch.matmul(x,torch.matmul(A, torch.matmul(A.T, x.T)))))))
用户已能实现简单的xAx^T运算,对应的einsum语句为:
torch.einsum('bi,ij,bj -> b',B,A,B)
解答
首先拆解目标表达式:x@A^T@A@x^T@x@A@A^T@x^T 本质是两个样本级标量的乘积:
term1 = x@A^T@A@x^T:每个样本对应的标量结果term2 = x@A@A^T@x^T:每个样本对应的标量结果
分步实现(更直观)
通过einsum分别计算两个标量后相乘,最后调整形状为(m,1):
# 计算 term1 = x@A^T@A@x^T term1 = torch.einsum('bi,ji,jk,bk -> b', B, A, A, B) # 计算 term2 = x@A@A^T@x^T term2 = torch.einsum('bi,ij,kj,bk -> b', B, A, A, B) # 合并结果并调整形状 result = (term1 * term2).unsqueeze(-1)
一步到位的einsum语句
直接通过一次einsum完成全部运算,再调整形状:
result = torch.einsum('bi,ji,jk,bk,bl,lm,nm,bn -> b', B, A, A, B, B, A, A, B).unsqueeze(-1)
验证正确性
将einsum实现与循环实现的结果对比,二者在浮点数精度范围内完全一致:
# 循环实现的结果 loop_result = result # einsum实现的结果 einsum_result = torch.einsum('bi,ji,jk,bk,bl,lm,nm,bn -> b', B, A, A, B, B, A, A, B).unsqueeze(-1) # 验证误差 print(torch.allclose(loop_result, einsum_result)) # 输出 True
内容的提问来源于stack exchange,提问作者SimpleStudent
相关产品推荐
相关产品推荐

