如何加速Numpy中查找最优索引的计算?
如何加速Numpy中查找最优索引的计算?
我有一个Numpy数组,用来把x-y坐标映射到对应的z坐标。具体来说,我用了一个2D数组,x和y作为它的两个轴,数组里存的是对应的z值:
import numpy as np x_size = 2000 y_size = 2500 z_size = 400 rng = np.random.default_rng(123) z_coordinates = np.linspace(0, z_size, y_size) + rng.laplace(0, 1, (x_size, y_size))
每个2000*2500的x-y点都对应一个z值(0到400之间的浮点数)。现在我想为每个整数z和整数x,找出最匹配的y值——本质上就是要创建一个形状为(x_size, z_size)的映射数组,里面存的是对应的最优y值。
最直接的思路是先创建一个目标形状的空数组,然后遍历每个z值:
y_coordinates = np.empty((x_size, z_size), dtype=np.uint16) for i in range(z_size): y_coordinates[:, i] = np.argmin( np.abs(z_coordinates - i), axis=1, )
但这个方法在我的机器上要跑11秒左右,速度实在慢得难以接受。
我当然试过更向量化的方法,理论上应该更快,比如:
y_coordinates = np.argmin( np.abs( z_coordinates[..., np.newaxis] - np.arange(z_size) ), axis=1, )
但出乎意料的是,这个版本比上面的循环还要慢60%左右(我用1/10的规模测试的,因为全量运行会占用巨量内存)。
另外,我还试过用numba的@jit(nopython=True)装饰器把代码包装成函数,结果也没起到加速效果。
请问怎么才能加速这个计算过程?
备注:内容来源于stack exchange,提问作者YPOC
相关产品推荐
相关产品推荐

