【问题标题】:How to improve the performance of this tiny distance Python function如何提高这个微小距离 Python 函数的性能
【发布时间】:2014-09-12 08:33:37
【问题描述】:

在为sklearn 的聚类算法使用自定义距离度量函数时,我遇到了性能瓶颈。

Run Snake Run 显示的结果是这样的:

很明显,问题出在dbscan_metric 函数上。该函数看起来很简单,我不太清楚加速它的最佳方法是什么:

def dbscan_metric(a,b):
  if a.shape[0] != NUM_FEATURES:
    return np.linalg.norm(a-b)
  else:
    return np.linalg.norm(np.multiply(FTR_WEIGHTS, (a-b)))

任何关于是什么导致它如此缓慢的想法将不胜感激。

【问题讨论】:

  • 这些数组有多大?如果您摆脱 if 语句并为数据集执行两个语句之一,这会加快速度吗? ...您可以尝试len(a) != NUM_FEATURES,看看是否更快...
  • 答案取决于数组 a 和 b 的大小。您还使用哪个 numpy 版本?考虑到配置文件它们可能很小并且你被 python 开销所支配,在这种情况下你需要使用 cython 来减少它
  • 在 norm@linalg.py 中花费了 71 秒,而在其他地方花费了 170 秒,这看起来很奇怪吗?我以为我知道如何理解这张图,但这似乎很奇怪。我所能猜测的是,额外的 170 秒以某种方式涉及调用开销。你能试试内联吗?
  • NUM_FEATURES 和 FTR_WEIGHTS 的本地范围外查找也可能需要一段时间
  • @JoranBeasley:我也尝试过 len(a) 并且很相似。数组不是那么大; NUM_FEATURES 是 130。我相信我的整个数据集都是这样,但出于某种原因,sklearn 有时会调用长度较小的函数,这就是我必须添加长度检查的原因。

标签: python optimization numpy scikit-learn


【解决方案1】:

我不熟悉该函数的作用 - 但是否有可能重复计算?如果是这样,您可以记住该功能:

cache = {}
def dbscan_metric(a,b):

  diff = a - b

  if a.shape[0] != NUM_FEATURES:
    to_calc = diff
  else:
    to_calc = np.multiply(FTR_WEIGHTS, diff)

  if not cache.get(to_calc): cache[to_calc] = np.linalg.norm(to_calc)

  return cache[to_calc]

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2012-01-13
    • 1970-01-01
    • 1970-01-01
    • 2011-04-08
    • 2015-07-14
    • 2014-07-17
    相关资源
    最近更新 更多