Numpy高效计算B:=Σv[j]·h(u[j]·A)的最优实现方法咨询
设h为逐元素作用的函数(例如平方运算等),u、v为长度为k的一维数组,A为形状为n×m的二维数组,其中n、m的取值可以非常大。
问题:在numpy中,如何高效计算如下定义的n×m数组B:
B := v[0] * h(u[0] * A) + ... + v[k-1] * h(u[k-1] * A)?
已知观察结论
- 朴素解法可直接遍历所有u[j]、v[j],累加
v[j] * h(u[j] * A)的计算结果。 - 另一种至少在内存层面表现极差的实现写法为:
B = np.sum(v[:, None, None] * h(u[:, None, None] * A), axis=0)
该写法会先生成k个与A尺寸完全一致的临时数组,内存占用随k线性增长,n、m较大时会直接触发内存不足错误。
最优实现方案
当n、m取值很大时,内存稳定性优先级远高于消灭Python层循环的开销,直接采用遍历k的朴素写法就是最优解:
import numpy as np # 可替换为任意逐元素运算的h函数 def h(x): return x ** 2 B = np.zeros_like(A) for u_j, v_j in zip(u, v): B += v_j * h(u_j * A)
方案优势
- 内存占用极低:全程仅需要2个和A同尺寸的数组(原数组A、结果数组B),单次迭代产生的临时数组计算完立即回收,内存占用不会随k的增大而升高
- 计算效率足够高:每次迭代的乘法、h函数运算、累加都是numpy原生向量化操作,Python层循环k次的开销可以完全忽略,计算速度和广播sum的写法几乎无差异
内容的提问来源于stack exchange,提问作者dohmatob
相关产品推荐
相关产品推荐

