如何在NumPy中高效比较三维数组与二维数组?
我有一个形状为(height, width, 3)的三维数组,代表一张BGR图像,数组值为[0,1]区间内的浮点数。对像素进行操作后,得到一个形状为(height, width)的二维数组,数组值是每个像素操作后的结果。
现在我需要将原始图像数组与这个结果数组进行比较,具体来说,要比较每个像素的BGR分量与结果数组对应坐标的值。例如,我先通过以下代码得到每个像素的最大BGR分量值:
import numpy as np img = np.random.random((360, 640, 3)) maxa = img.max(axis=-1)
直接使用img == maxa会报错:
In [335]: img == maxa <ipython-input-335-acb909814b9a>:1: DeprecationWarning: elementwise comparison failed; this will raise an error in the future. img == maxa Out[335]: False
我用Python嵌套循环实现了预期逻辑,但效率极低:
result = [[[c == maxa[y, x] for c in img[y, x]] for x in range(640)] for y in range(360)]
之后用img == np.dstack([maxa, maxa, maxa])实现了相同功能,且效率提升明显,经测试结果正确:
In [339]: result = [[[c == maxa[y, x] for c in img[y, x]] for x in range(640)] for y in range(360)] ...: np.array_equal(arr3, img == np.dstack([maxa, maxa, maxa])) Out[339]: True
各方法的性能测试结果如下:
In [340]: %timeit [[[c == maxa[y, x] for c in img[y, x]] for x in range(640)] for y in range(360)] 509 ms ± 16.6 ms per loop (mean ± std. dev. of 7 runs, 1 loop each) In [341]: maxals = maxa.tolist() In [342]: imgls = img.tolist() In [343]: %timeit [[[c == maxals[y][x] for c in imgls[y][x]] for x in range(640)] for y in range(360)] 156 ms ± 2.57 ms per loop (mean ± std. dev. of 7 runs, 10 loops each) In [344]: %timeit img == np.dstack([maxa, maxa, maxa]) 4.25 ms ± 121 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
请问有没有更高效的实现方法?
可以利用NumPy的广播机制来避免显式堆叠数组,这会比np.dstack更高效,因为不需要额外创建堆叠后的中间数组。
方案1:扩展维度后直接比较
通过给maxa增加一个维度,让它的形状变为(height, width, 1),这样NumPy会自动将其广播到与img匹配的(height, width, 3)形状,直接进行元素级比较:
# 写法1:使用np.newaxis扩展维度 result = img == maxa[..., np.newaxis] # 写法2:更简洁的None语法(np.newaxis的别名) result = img == maxa[..., None]
方案2:使用np.expand_dims扩展维度
如果觉得索引语法不够直观,也可以用np.expand_dims来显式扩展维度:
result = img == np.expand_dims(maxa, axis=-1)
性能对比
测试上述方法的性能,会发现比np.dstack更快:
%timeit img == maxa[..., None] # 输出示例:2.1 ms ± 45 µs per loop (mean ± std. dev. of 7 runs, 100 loops each) %timeit img == np.expand_dims(maxa, axis=-1) # 输出示例:2.08 ms ± 32 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
原理说明
NumPy的广播机制会自动处理形状兼容的数组运算:当maxa的形状是(360,640),扩展为(360,640,1)后,与img的(360,640,3)进行比较时,NumPy会将maxa的每个值沿最后一个维度“虚拟复制”3次,完成元素级匹配——这个过程是在底层优化的,不需要显式创建复制后的数组,因此比np.dstack更节省内存和运算时间。
内容的提问来源于stack exchange,提问作者Ξένη Γήινος

