You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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 本质是两个样本级标量的乘积:

  1. term1 = x@A^T@A@x^T:每个样本对应的标量结果
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 05:27:04