无额外库依赖实现NumPy数组第一轴并行循环 替代显式for循环
假设现有如下NumPy数组:
prova = np.array([[[0, 8], [8, 8], [7, 7]], [[6, 5], [6, 6], [0, 3]], [[0, 1], [5, 6], [0, 2]], [[6, 2], [2, 8], [4, 4]]])
另有一个示例函数,接收2个NumPy数组作为输入(函数具体逻辑仅作随机示例),始终返回int类型标量值:
def random_function(a1, a2): a1 = a1[np.argsort(a1)] a2 = a2[np.argsort(a2)] if a1[0]>a2[0]: print("1") return np.diff(a2-a1).max() elif a2[0]>a1[0]: print("2") return np.diff(a2-a1).min() else: print("3") return (np.cumsum(a1)+np.cumsum(a2))[-1]
需求是沿prova数组的第一轴应用该函数,原有基于for循环的实现写法如下:
np.array([random_function(*p.T) for p in prova])
运行后输出结果为:
array([-6, -4, -1, 26])
待解决问题:
- 如何在不使用for循环的前提下得到相同结果?
- 是否可以使用
np.apply_along_axis或np.apply_over_axes实现该目标?
首先明确两个apply类函数的适配性:np.apply_along_axis和np.apply_over_axes本质没有脱离Python层循环,只是把循环逻辑封装在了NumPy内部,性能和手写列表推导式没有本质区别,甚至会因为额外的参数校验开销更慢。同时这两个函数设计上是单个数组沿指定轴传参,没法直接适配当前场景下每次传入两个子数组的调用形式,硬写适配逻辑会非常冗余,可读性差,不推荐使用。
如果只是想去掉代码里显式的for关键字,不追求性能提升,原有的列表推导式已经是这类带分支逻辑的自定义函数逐样本调用场景下,可读性最高、额外开销最小的写法,不需要额外修改。
如果需要真正消除Python层循环、获得NumPy向量化的性能提升,需要把原函数的逻辑改写为支持批量维度操作的版本,具体实现如下:
- 先拆分批量输入的两个子数组,替代原循环里的
p.T操作:
# 拆分后a1、a2形状均为(4,3),第一维对应原循环的4个迭代样本 a1 = prova[:, :, 0] a2 = prova[:, :, 1]
- 批量完成排序操作,替代原函数里的argsort索引取值:
# 沿最后一个维度计算排序索引,批量完成排序 idx1 = np.argsort(a1, axis=-1) idx2 = np.argsort(a2, axis=-1) a1_sorted = np.take_along_axis(a1, idx1, axis=-1) a2_sorted = np.take_along_axis(a2, idx2, axis=-1)
- 生成三个分支的掩码,替代原函数里的if条件判断:
mask1 = a1_sorted[:, 0] > a2_sorted[:, 0] # 对应a1首元素更大的分支 mask2 = a2_sorted[:, 0] > a1_sorted[:, 0] # 对应a2首元素更大的分支 mask3 = ~(mask1 | mask2) # 对应两者首元素相等的分支
- 分别计算三个分支的结果,按掩码合并得到最终输出:
res = np.empty(a1.shape[0], dtype=np.int64) diff = a2_sorted - a1_sorted diff_pad = np.diff(diff, axis=-1) # 填充三个分支的计算结果 res[mask1] = diff_pad[mask1].max(axis=-1) res[mask2] = diff_pad[mask2].min(axis=-1) cumsum_total = np.cumsum(a1_sorted[mask3], axis=-1) + np.cumsum(a2_sorted[mask3], axis=-1) res[mask3] = cumsum_total[:, -1]
最终得到的res和原循环输出完全一致,为array([-6, -4, -1, 26])。该实现全部使用NumPy原生的C级向量化操作,没有Python层循环,在数组规模较大时性能会远高于循环、apply类、np.vectorize等实现。
注:原函数里的print输出属于副作用,向量化实现不会逐样本打印"1""2""3",如果需要保留打印逻辑,就无法完全消除Python层循环。
内容的提问来源于stack exchange,提问作者Salvatore Daniele Bianco

