使用Numba JIT优化K近邻函数反而变慢2倍,求问题排查
Numba优化K近邻函数反而变慢的问题分析
尝试用Numba JIT优化自定义的K近邻(KNN)向量求解函数,结果Numba版本的运行速度比原代码慢两倍(相同输入下:100000个训练点、10000个测试点,Numba耗时约260秒,原代码约130秒)。原代码和Numba优化代码如下:
原代码
def __call__(self, x_test: list[list] | np.ndarray) -> np.ndarray: x_test = np.asarray(x_test) y_test = np.empty(x_test.shape[0], dtype=self._y.dtype) for i in range(x_test.shape[0]): distances = np.linalg.norm(self._x - x_test[i], ord=self.ord, axis=1) label, count = np.unique(self._y[np.argsort(distances)[:min(self.k, self._y.shape[0])]], return_counts=True) y_test[i] = label[np.argmax(count)] return y_test
Numba优化代码
@overload(np.unique) def np_unique(a): def np_unique_impl(a): b = np.sort(a.flatten()) unique = list(b[:1]) counts = [1 for _ in unique] for x in b[1:]: if x != unique[-1]: unique.append(x) counts.append(1) else: counts[-1] += 1 return np.asarray(unique), np.asarray(counts) return np_unique_impl @jit(nopython=True) def jitter(x_train: np.ndarray, y_train: np.ndarray, k: int, ord: int, x_test: np.ndarray, y_test: np.ndarray): for i in range(x_test.shape[0]): distances = np.zeros(x_train.shape[0]) for j in range(distances.shape[0]): distances[j] = np.linalg.norm(x_train[j]-x_test[i], ord=ord) label, count = np.unique(y_train[np.argsort(distances)[:min(k, y_train.shape[0])]]) y_test[i] = label[np.argmax(count)]
问题根源分析
1. 距离计算的低效实现
原代码利用Numpy的广播机制(self._x - x_test[i])批量计算所有训练点与测试点的距离,底层是高度优化的C实现,能充分利用CPU缓存和指令级并行。而Numba版本改用Python循环逐个计算每个训练点的距离,完全抛弃了矢量化的优势,反而比Numpy原生实现慢。
2. 自定义np.unique的性能瓶颈
你重载的np.unique用Python列表来累加唯一标签和计数,在Numba的nopython模式下,列表操作的效率远低于Numpy原生的np.unique(原生是经过极致优化的C代码)。此外,额外的flatten和sort操作也增加了不必要的计算开销。
3. 全排序的冗余开销
原代码用np.argsort对所有距离排序后取前k个,但实际上不需要全排序——np.argpartition可以在O(n)时间内找到前k小的元素,比O(n log n)的全排序效率高得多,尤其是当训练集规模很大时。Numba版本沿用了全排序,进一步拖慢了速度。
优化后的Numba实现方案
针对上述问题,优化后的代码保留Numpy矢量化优势,替换低效操作,同时开启并行:
from numba import jit, prange import numpy as np @jit(nopython=True, parallel=True) def optimized_knn(x_train, y_train, k, ord, x_test): n_test = x_test.shape[0] n_train = x_train.shape[0] y_test = np.empty(n_test, dtype=y_train.dtype) k_eff = min(k, n_train) # 假设标签为非负整数,若标签是其他类型可先做映射 max_label = y_train.max() counts = np.zeros(max_label + 1, dtype=np.int64) for i in prange(n_test): # 保留矢量化距离计算,Numba会自动优化 distances = np.linalg.norm(x_train - x_test[i], ord=ord, axis=1) # 用argpartition快速找到前k小的距离索引,替代全排序 top_k_idx = np.argpartition(distances, k_eff)[:k_eff] top_k_labels = y_train[top_k_idx] # 重置计数数组,统计top-k标签的出现次数 counts[:] = 0 for label in top_k_labels: counts[label] += 1 # 取出现次数最多的标签 y_test[i] = np.argmax(counts) return y_test
优化点说明
- 矢量化距离计算:保留Numpy的广播式距离计算,Numba能很好地优化这类操作,效率接近原生Numpy甚至更高。
argpartition替代全排序:将排序开销从O(n log n)降到O(n),大训练集下提升明显。- 计数数组统计众数:避免
np.unique的额外开销,直接用数组计数,效率更高。 - 并行处理测试点:开启
parallel=True并使用prange,利用多核CPU并行处理每个测试点,大幅缩短总耗时。
内容的提问来源于stack exchange,提问作者Anili
相关产品推荐
相关产品推荐

