【问题标题】:How to optimize a numpy loop that sums values from an array which is indexed by another array where values equal the loop index如何优化一个numpy循环,该循环对数组中的值求和,该数组由另一个数组索引,其中值等于循环索引
【发布时间】:2018-03-30 08:37:29
【问题描述】:

我有这段代码在应用程序运行期间被多次调用。 它需要一个代表值的数字数组(value_array)。 这些应该在 zone_array 中定义的区域中总结。 zone_ids 表示 zone_array 中所有可能区域的列表。

它基本上是这样的:我有一张人口栅格地图,我想知道有多少人住在区域地图的每个区域。

代码:

values = np.zeros(len(zone_ids))
for i in zone_ids:
    values[i] = round(np.nansum(value_array[zone_array == i]), 2)
return values

罪魁祸首似乎是for循环,但我还没有找到消除它的方法并得到相同的结果。

我用 bincount 尝试过,但没有成功。 使用 numba jit 也没有效果。

我想远离 cython,因为此代码将用于不支持 cython 的 Qgis 插件。

测试代码:

import numpy as np


def fill_values(zone_array, value_array, zone_ids):
    values = np.zeros(len(zone_ids))
    for i in zone_ids:
        values[i] = round(np.nansum(value_array[zone_array == i]), 2)
    return values


def run():
    # 300 different zones
    zone_ids = range(300)
    # zone map with 300 zones
    zone_array = (np.random.rand(2000, 2000) * 300).astype(int)
    # value map from which we want the sum of values per zone (real map can have NaN values)
    value_array = (np.random.rand(2000, 2000) * 10.)
    value_array[5, 5] = np.NAN
    fill_values(zone_array, value_array, zone_ids)


if __name__ == '__main__':
    run()

每个循环 1.92 秒 ± 17.5 毫秒(平均值 ± 标准偏差,7 次运行,每个循环 1 个)

按照 Divakar 的建议实施 bincount:

每个循环 203 毫秒 ± 15.2 毫秒(7 次运行的平均值 ± 标准偏差,每个循环 1 个)

【问题讨论】:

  • 罪魁祸首不是for循环。相反,问题在于内部的比较zone_array==i。对于每个 zone_id i,必须检查所有 2000x2000=4e6 值是否与 i 相等。
  • 如果我减少区域 id 的数量,我的速度会提高,所以 for 循环仍然涉及性能问题。而且由于我没有其他选择,我知道不做zone_array==i 我专注于循环。最好的办法是我可以以某种方式使用zone_array == zone_ids 并跳过循环。
  • 您可以使用zone_array[:,:,None] == zone_ids 广播比较结果,但这仍然会在 for 循环中留下索引,并且对性能没有太大的改进。

标签: python performance numpy for-loop


【解决方案1】:

直接使用bincount,您将在求和中得到NaNs。因此,您可以简单地将NaNs 替换为zeros 并使用bincount。作为矢量化解决方案,这应该更快。

因此,实现将是 -

val_nonan = np.where(np.isnan(value_array), 0, value_array)
out = np.round(np.bincount(zone_array.ravel(), val_nonan.ravel()),2)

【讨论】:

  • 这适用于我的问题。非常感谢。我猜我的 bincount 会尝试被 nan 值弄乱的地方。另外values = out[zone_ids] 用于您想要区域子集结果的情况。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2020-11-20
  • 1970-01-01
  • 1970-01-01
  • 2018-05-09
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多