【发布时间】: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