用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__
相关产品推荐
相关产品推荐

