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

用EinSum替代Reshape实现更快的NumPy计算

问题描述

在Python中,我们有维度为(T,)的数组A、(L,T)的数组U、(K,T)的数组G以及(L,L,T)的数组Y。当前代码可生成维度为(T, LK, 1)的numer1和numer2,以及维度为(LK,1)的numer3。但实际运行中np.reshape速度极慢,甚至慢于循环实现。能否通过调整EinSum的使用方式来避免reshape操作(毕竟最后要执行求和操作)?类似的优化也适用于以下分母计算代码,该场景下reshape的性能问题更为突出。

最小可复现示例(MWE)
np.random.seed(0)
T=250
K = 15
L = 20
A = np.random.normal(size=(T))
Y = np.random.normal(size=(L,L,T))
G = np.random.normal(size=(K,T))
U = np.random.normal(size=(L,T))
原计算代码
# Calculations
L, T = U.shape
K = G.shape[0]
Y_transposed = Y.transpose(2, 0, 1)
sG = np.einsum('it,jt->tij', G, G)
A_transposed = A[None, :, None].transpose(1, 0, 2)
numer1 = np.einsum('ai,ak->aik', U.T, G.T).reshape(T, L * K, 1)
numer2 = numer1 * A_transposed
numer3 = numer2.sum(axis=0)
原分母计算代码
denom1 = np.einsum('aij,akl->aikjl', Y_transposed, sG).reshape(T, L * K, L *K)
denom2 = denom1 * A_transposed
denom3 = denom2.sum(axis=0)
优化方案

完全可以通过调整einsum的下标定义,直接合并reshape和后续的乘法、求和操作,避免生成中间超大维度数组,从根源解决reshape的性能问题。

分子部分优化

原逻辑是先计算U.T与G.T的外积,reshape后乘以A的转置再求和。我们可以用einsum直接完成所有步骤,跳过中间的reshape:

# 直接生成最终的numer3,无需中间变量
numer3 = np.einsum('ti, tk, t -> ik', U, G, A).reshape(-1, 1)

逻辑说明:

  • U是(L,T),用下标ti表示时间步t对应的U第i个元素
  • G是(K,T),用下标tk表示时间步t对应的G第k个元素
  • A是(T,),用下标t表示时间步的权重
  • 对每个时间步t计算U[:,t] * G[:,t] * A[t],再对所有t求和,最后通过reshape(-1,1)转成(LK,1)的目标格式

分母部分优化

原分母代码会生成(T,L,K,L,K)的超大中间数组,reshape操作的性能损耗尤其严重。通过einsum合并步骤可彻底避免这个问题:

# 直接生成最终的denom3,无需中间变量
denom3 = np.einsum('tij, tkl, t -> (ik)(jl)', Y_transposed, sG, A)

逻辑说明:

  • Y_transposed是(T,L,L),用下标tij表示时间步t的(i,j)元素
  • sG是(T,K,K),用下标tkl表示时间步t的(k,l)元素
  • A是(T,),用下标t表示时间步的权重
  • 对每个时间步t计算Y_transposed[t,i,j] * sG[t,k,l] * A[t],再对所有t求和,最后通过(ik)(jl)的下标直接将二维组合下标合并为一维,得到(LK,LK)的目标格式,完全跳过了中间数组的生成与reshape操作

性能优势

优化后的代码不仅避免了reshape的性能损耗,还大幅减少了内存占用(尤其是分母部分),运行速度会远快于原代码,甚至超过循环实现的效率。

内容的提问来源于stack exchange,提问作者user9875321__

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 23:33:15