【问题标题】:Speed up Python nested loops with conditional statements使用条件语句加速 Python 嵌套循环
【发布时间】:2015-08-02 05:43:42
【问题描述】:

我正在将代码从 MATLAB 转换为 python,以加快简单的操作。我编写了一个包含嵌套循环和条件语句的函数;循环的目的是返回与数组 y 相比时数组 x 中最近元素的索引列表。我按 1e5 个项目的顺序进行比较,运行时间约为 30 秒。任何有助于加快此过程的帮助将不胜感激!我在使用 numba-pro 自动即时编译器方面取得了部分成功:

@autojit()
def find_nearest(x,y,idx):
    idx_old = 0
    rng1 = range(y.shape[0])
    rng2 = range(x.shape[0])
    for i in rng1:
        prev = abs(x[idx_old]-y[i])
        for j in rng2:
            if abs(x[j]-y[i]) < prev:
                prev = abs(x[j]-y[i])
                idx_old = j
        idx[i] = idx_old
    return idx

对不起,我是个菜鸟,我是python的新手!

【问题讨论】:

  • 您能否更新您的脚本以包含find_nearest 的示例数据输入,这样就清楚了吗?
  • 例如 x = np.array([1.1,2.3,5.9,8.5]), y = np.array([0.2, 5.5, 12]) 和 idx = np.zeros(np. shape(y)) 应该返回 idx = [0,2,3] (x 中最接近 y 中的项目的索引。在 MATLAB 中,我使用 knnsearch 执行此操作,在我的大型数据集上大约需要 2.5 秒解决;而我的实现大约需要 30 秒。输入数组不需要按任何特定的排序顺序。谢谢你看!
  • 我尝试使用 sci-kitlearn 的 k-nearest-neighbors 实现,但是它返回了一个在数据集上训练的函数;并且对我的完整数据集进行培训是不可行的。我的搜索域包含 1023848 个项目,我正在尝试在其中找到 12325 个最接近的项目。

标签: performance python-2.7 numpy conditional-statements numba


【解决方案1】:

您的 Numba 代码没有任何问题,只是算法效率不高。更好的是对x 数组进行排序并进行二分查找,非常类似于this answer 和this answer:

def find_nearest(x, y):
    indices = np.argsort(x)

    loc = np.searchsorted(x[indices], y)
    right = indices.take(loc, mode='clip')
    left = indices.take(loc-1, mode='clip')

    return np.where(abs(y-x[left]) < abs(y-x[right]), left, right)

在我的 PC 上,这比 x 和 y 分别具有 106 和 105 元素的 KDTree 方法快了大约 80 倍。大约三分之二的时间都花在了argsort-ing 数组上,所以我认为在这里使用 Numba 不会有太多收获。

【讨论】:

  • 非常感谢您的回复;我能够实现你的代码。它在我的数据集上运行速度提高了大约 270 倍(在一组大约 1e6 个项目中找到大约 1e4 个索引)。您能否推荐一个您发现在学习用 Python 实现更高效的代码时有用的教程/文章/网页?我查看了您建议的类似答案。索引操作总是那么快吗?
  • @Chris 我主要通过关注本网站上的[Numpy] 标记和试验/试错来学习编写更高效的 Python。您能否进一步解释一下“索引操作总是这么快”是什么意思?
  • 通过索引操作,我指的是返回布尔/索引数组以进行后续操作的函数。我最初试图将我的“find_nearest”函数变成一行矩阵和向量运算;然而,当我尝试在大量数据集上运行它时遇到了“MemoryError”。
  • @Chris 啊,我明白了。不,这与索引操作与矩阵/向量操作无关。您的原始算法的复杂度为 O(m*n),其中m 和n 是x 和y 的长度。您可能还创建了一个大小为m*n(非常大)的数组。使用我的代码,复杂度为 O((m+n)log(m)),额外的内存使用量为 O(m+n)(我认为.. 不确定排序)。
【解决方案2】:

我找到了一个临时解决方案来解决我的问题。通过实现 scipy.spatial 的 kdtree,我能够将运行时间从 32 秒缩短到不到 10 秒。这仍然比 MATLAB knnsearch 算法慢四倍;了解如何使用条件语句加速循环仍然很重要。但目前这个修改后的实现更快:

from scipy import spatial
from numpy import matrix

tree = spatial.KDTree(matrix(x).T)
(_, idxx) = tree.query(matrix(y).T)

数组 x 和 y 是平面 1d 格式;树要求查询采用列向量形式。

任何改进原始实现的运行时间的建议将不胜感激!

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2017-12-20
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2019-04-22
    • 2013-11-01
    • 2021-04-27
    相关资源
    最近更新 更多