You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 08:28:15