从零实现KNN算法:原理、Python代码与手写数字识别实战
1. 项目概述:从零理解KNN算法
KNN,全称K-Nearest Neighbors,中文常译为K近邻算法。我第一次接触它,是在一个手写数字识别的项目里,当时觉得这算法简直“简单粗暴”得不像个正经的机器学习方法。它不像那些复杂的神经网络,需要你费尽心思去设计层数和激活函数,也不像支持向量机那样有深厚的数学理论支撑。KNN的核心思想就一句话:物以类聚,人以群分。一个新来的数据点,看看它周围离得最近的K个邻居都是谁,它大概率就跟这些邻居属于同一类。这种基于实例的学习,或者说“懒惰学习”,让它在很多场景下成为了一个快速验证想法、建立基线模型的绝佳工具。
对于刚入门机器学习的朋友来说,KNN是一个完美的起点。它几乎不需要你理解复杂的数学推导,直观易懂,而且用Python实现起来代码量极少,能让你快速获得“我搞定了机器学习”的正向反馈。无论是想识别图片里的数字,还是根据用户特征进行简单的分类推荐,KNN都能提供一个可靠的基线方案。当然,它也有自己的局限,比如计算开销大、对数据尺度敏感等,这些我们后面会详细拆解。今天,我们就从最根本的原理出发,手把手用Python实现一个“五脏俱全”的KNN分类器,并探讨如何让它真正用起来,而不是仅仅停留在“Hello World”的演示阶段。
2. KNN算法核心原理与设计思路拆解
2.1 “物以类聚”的数学表达:距离度量
KNN算法的所有魔力,都建立在“距离”这个概念之上。我们说一个新样本的类别由其最近的K个邻居决定,那么如何定义“最近”?这就需要距离度量。最常用的是欧氏距离,也就是我们中学学过的两点间直线距离。在二维空间里,点 (x1, y1) 和点 (x2, y2) 的欧氏距离是 √[(x1-x2)² + (y1-y2)²]。推广到n维特征空间,公式也类似。
但欧氏距离不是唯一的选择。在文本分类或者某些特定场景下,曼哈顿距离(各维度坐标差绝对值的和)可能更合适,它计算的是沿着坐标轴行走的“城市街区”距离。还有一种叫闵可夫斯基距离,算是欧氏距离和曼哈顿距离的通用形式。选择哪种距离度量,没有绝对的金科玉律,它取决于你的数据特性和业务逻辑。比如,如果你的特征向量是稀疏的(很多0值),余弦相似度(计算两个向量夹角的余弦值)可能比欧氏距离更能反映其相似性。
注意:距离度量的选择直接影响邻居的选取,从而决定分类结果。在实际操作中,如果特征量纲不统一(比如一个特征是“年薪(万元)”,另一个特征是“年龄”),直接计算欧氏距离会被大数值特征主导,导致距离失真。因此,数据标准化(如Z-score标准化)或归一化(缩放到[0,1]区间)是使用KNN前几乎必不可少的预处理步骤。很多新手会忽略这一点,导致模型效果莫名其妙地差。
2.2 关键参数K的选择:平衡的艺术
K值是这个算法中最重要的超参数,没有之一。K太小(比如K=1),模型会变得非常敏感,容易受到噪声数据或异常点的干扰,导致模型过拟合,即在训练集上表现很好,但在新数据上表现糟糕。想象一下,你家隔壁搬来一个行为古怪的新邻居,如果只根据他一个人来判断整条街的风气,结论很可能有失偏颇。
反之,如果K值取得太大,模型又会变得过于“平滑”或“懒惰”。它会考虑太多远处的点,以至于决策边界变得模糊,可能无法捕捉到数据中细微的、局部的模式,导致欠拟合。这就好比要通过全市人民的平均意见来决定你家小区是否该修个花园,显然忽略了本小区的实际需求。
那么,如何选择K呢?一个最经典且实用的方法是交叉验证。我们可以把训练数据分成多份,用其中一部分训练,另一部分验证,尝试不同的K值(通常从1开始,取到训练样本数的平方根左右),选择在验证集上平均准确率最高的那个K。在实现时,我们通常会选择奇数的K值,以避免在二分类问题中出现平票的尴尬局面。
2.3 决策规则:邻居们如何投票?
找到了K个最近的邻居后,如何根据他们来决定新样本的类别?最常用的方法是多数表决。也就是看这K个邻居中,哪个类别的样本数最多,新样本就属于那个类别。这是最直观的方式。
但有时候,我们觉得距离更近的邻居应该拥有更大的话语权。这就引入了加权投票的方法。常见的权重设置是距离的倒数,即距离越近,权重越大。这样,一个紧挨着的邻居的一票,可能抵得上远处三个邻居的票。加权投票在处理类别分布不均匀或者噪声数据时,往往能获得更鲁棒的效果。
除了分类,KNN也可以用于回归任务。在KNN回归中,对于一个新的样本点,我们取其K个最近邻居的目标值的平均值(或加权平均值)作为该样本的预测值。这同样体现了“近朱者赤”的思想。
3. 从零实现KNN分类器的核心细节
3.1 数据结构与算法流程设计
在动手写代码之前,我们先在脑子里把流程过一遍。一个完整的KNN分类器,其工作流程可以清晰地分为两个阶段:训练阶段和预测阶段。
训练阶段出奇地简单:KNN是一种“懒惰学习”算法,它实际上并不从训练数据中学习一个显式的模型(比如一条直线或一个复杂的函数)。它的训练过程仅仅是将训练数据集(特征矩阵X_train和标签向量y_train)存储起来。因此,我们的fit方法可能只有一行代码:self.X_train = X_train; self.y_train = y_train。这也是为什么KNN训练速度“极快”的原因——它几乎什么都没做。
真正的计算发生在预测阶段。对于一个待预测的新样本,我们需要:
- 计算该样本与训练集中每一个样本的距离。
- 从这些距离中,找出最小的K个(即最近的K个邻居)。
- 查看这K个邻居对应的标签。
- 根据投票规则(如多数表决)确定新样本的预测标签。
这个流程决定了KNN预测速度慢的缺点,因为每次预测都需要与所有训练样本计算距离。当训练集很大时,计算开销会变得难以承受。这也是后续优化(如KD树、球树)要解决的核心问题。
3.2 距离计算的高效实现
计算距离是KNN中最耗时的部分。我们需要高效地计算一个样本与所有训练样本的距离。这里可以利用NumPy的广播机制进行向量化运算,避免低效的Python循环。
假设我们的训练集self.X_train是一个形状为(n_samples_train, n_features)的矩阵,待预测的单个样本x是一个形状为(n_features,)的向量。计算欧氏距离的平方(为了避免开方运算,节省时间,因为开方不影响大小顺序)可以这样向量化实现:
import numpy as np # 计算差值 diff = self.X_train - x # 广播发生,得到 (n_samples_train, n_features) 的矩阵 # 计算平方和 distances = np.sum(diff ** 2, axis=1) # 沿特征轴求和,得到 (n_samples_train,) 的距离平方向量这段代码一次性计算了所有距离,效率远高于写一个for循环。对于曼哈顿距离,只需将diff ** 2改为np.abs(diff)即可。
3.3 邻居选取与投票机制
得到所有距离后,我们需要找到最小的K个距离对应的索引。np.argsort函数可以对数组排序并返回索引,但我们只需要前K个最小的,使用np.argpartition函数会更快,因为它只进行部分排序。
# 获取距离最小的K个样本的索引 k_nearest_indices = np.argpartition(distances, kth=self.k)[:self.k] # 根据索引获取这K个邻居的标签 k_nearest_labels = self.y_train[k_nearest_indices]接下来是投票。对于多数表决,我们可以使用np.bincount来统计每个标签出现的次数,然后取最大值对应的标签。但np.bincount要求标签是非负整数。如果我们的标签是字符串或其他类型,可以先用np.unique映射一下,或者直接用collections.Counter。
from collections import Counter # 使用Counter统计并找出最常见的标签 most_common_label = Counter(k_nearest_labels).most_common(1)[0][0]如果要实现加权投票,过程会稍微复杂一些。我们需要根据距离计算权重(例如weights = 1.0 / (distances[k_nearest_indices] + 1e-5),加一个极小值防止除零),然后为每个类别累加其邻居的权重,最后取权重和最大的类别作为预测结果。
4. 手写KNN分类器的完整Python实现
下面,我们将上述思路整合,实现一个功能完整的KNN分类器类。这个实现将包含核心的fit和predict方法,并考虑一些工程细节。
4.1 类结构设计与初始化
我们首先定义KNNClassifier类。在__init__方法中,我们主要接收超参数n_neighbors(即K值),并可以预留一个参数用于选择距离度量方式(这里我们先实现欧氏距离)。
import numpy as np from collections import Counter from sklearn.base import BaseEstimator, ClassifierMixin # 可选,用于兼容scikit-learn API class KNNClassifier: """ 一个从头实现的K近邻分类器。 参数: n_neighbors (int): 用于投票的邻居数量K。 weights (str): 投票权重。'uniform'为等权投票,'distance'为距离倒数加权。 metric (str): 距离度量。'euclidean'为欧氏距离,'manhattan'为曼哈顿距离。 """ def __init__(self, n_neighbors=5, weights='uniform', metric='euclidean'): self.n_neighbors = n_neighbors self.weights = weights self.metric = metric self.X_train = None self.y_train = None def fit(self, X, y): """ 训练KNN模型。实际上只是存储训练数据。 参数: X (np.ndarray): 训练特征,形状 (n_samples, n_features)。 y (np.ndarray): 训练标签,形状 (n_samples,)。 返回: self: 返回实例本身。 """ # 简单的输入检查 if X.shape[0] != y.shape[0]: raise ValueError("训练样本数必须与标签数一致。") self.X_train = np.array(X) self.y_train = np.array(y) return self这里我们遵循了scikit-learn的API设计惯例(fit返回self),这样我们的模型以后可以更方便地嵌入到scikit-learn的管道(Pipeline)中。输入检查是保证代码健壮性的好习惯。
4.2 核心预测方法的实现
predict方法需要能够处理单个样本和批量样本。我们实现一个内部的_predict_one方法来处理单个样本的预测,然后在predict中通过循环或向量化方式处理批量输入。
def _predict_one(self, x): """ 预测单个样本的标签。 参数: x (np.ndarray): 单个样本的特征向量,形状 (n_features,)。 返回: predicted_label: 预测的标签。 """ # 1. 计算距离 if self.metric == 'euclidean': # 计算欧氏距离的平方 distances = np.sum((self.X_train - x) ** 2, axis=1) elif self.metric == 'manhattan': # 计算曼哈顿距离 distances = np.sum(np.abs(self.X_train - x), axis=1) else: raise ValueError(f"不支持的度量方式: {self.metric}") # 2. 获取K个最近邻居的索引 # 使用argpartition进行部分排序,比完全排序argsort更快 k_nearest_indices = np.argpartition(distances, kth=self.n_neighbors)[:self.n_neighbors] # 3. 获取K个邻居的标签 k_nearest_labels = self.y_train[k_nearest_indices] # 4. 投票决策 if self.weights == 'uniform': # 多数表决 most_common = Counter(k_nearest_labels).most_common(1) return most_common[0][0] elif self.weights == 'distance': # 加权投票(距离倒数作为权重) k_nearest_distances = distances[k_nearest_indices] # 防止距离为0导致除零错误,加一个极小值 weights = 1.0 / (k_nearest_distances + 1e-10) # 为每个类别累加权重 weight_dict = {} for label, weight in zip(k_nearest_labels, weights): weight_dict[label] = weight_dict.get(label, 0.0) + weight # 返回权重和最大的标签 return max(weight_dict.items(), key=lambda x: x[1])[0] else: raise ValueError(f"不支持的权重方式: {self.weights}") def predict(self, X): """ 预测批量样本的标签。 参数: X (np.ndarray): 待预测样本特征,形状 (n_samples, n_features)。 返回: predictions (np.ndarray): 预测标签数组,形状 (n_samples,)。 """ if self.X_train is None: raise ValueError("模型尚未训练,请先调用fit方法。") X = np.array(X) # 对每个样本应用_predict_one predictions = np.array([self._predict_one(x) for x in X]) return predictions这个实现已经具备了核心功能。_predict_one方法清晰地展示了KNN预测的四个步骤。在批量预测时,我们使用了列表推导式,对于非常大的数据集,可以考虑进一步向量化优化,但当前版本对于理解和教学来说已经足够清晰。
4.3 模型评估与K值选择
模型写好了,我们怎么知道它好不好?我们需要用数据来评估。通常我们会将数据集划分为训练集和测试集,用训练集来fit模型,用测试集来评估predict的准确率。
from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score # 假设 X, y 是你的特征和标签数据 X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) # 创建并训练模型 knn = KNNClassifier(n_neighbors=5) knn.fit(X_train, y_train) # 预测并评估 y_pred = knn.predict(X_test) accuracy = accuracy_score(y_test, y_pred) print(f"模型在测试集上的准确率为: {accuracy:.4f}")如何选择最优的K值?我们可以写一个简单的循环,尝试不同的K值,并用交叉验证来评估其性能,避免过拟合到某一次划分的数据上。
from sklearn.model_selection import cross_val_score import matplotlib.pyplot as plt # 尝试不同的K值 k_range = range(1, 31) k_scores = [] for k in k_range: knn = KNNClassifier(n_neighbors=k) # 使用5折交叉验证计算平均准确率 scores = cross_val_score(knn, X_train, y_train, cv=5, scoring='accuracy') k_scores.append(scores.mean()) # 绘制K值与准确率的关系图 plt.plot(k_range, k_scores) plt.xlabel('Value of K for KNN') plt.ylabel('Cross-Validated Accuracy') plt.show() # 找出最佳K值 best_k = k_range[np.argmax(k_scores)] print(f"交叉验证建议的最佳K值为: {best_k}")通过这个图,你可以清晰地看到模型性能随K值变化的趋势。通常,准确率会先随着K增大而提升(减少噪声影响),达到一个峰值后开始下降(模型过于平滑)。那个峰值对应的K值,往往就是比较理想的选择。
5. 实战演练:用KNN进行手写数字识别
理论学习之后,我们用一个经典的案例——手写数字识别,来检验我们的KNN分类器。这里我们使用scikit-learn内置的digits数据集,它包含了1797张8x8像素的手写数字图片。
5.1 数据加载与探索
首先,我们加载数据并看看它的样子。
from sklearn.datasets import load_digits import matplotlib.pyplot as plt digits = load_digits() X, y = digits.data, digits.target print(f"数据形状: {X.shape}") # (1797, 64) print(f"标签形状: {y.shape}") # (1797,) print(f"类别: {np.unique(y)}") # [0 1 2 3 4 5 6 7 8 9] # 可视化前10个数字 fig, axes = plt.subplots(2, 5, figsize=(10, 5)) for i, ax in enumerate(axes.flat): ax.imshow(X[i].reshape(8, 8), cmap='gray') ax.set_title(f"Label: {y[i]}") ax.axis('off') plt.show()你会发现,每个样本是一个64维的向量(将8x8的图片展平)。像素值范围是0-16。对于KNN来说,不同特征的量纲一致,所以我们可以暂时不做标准化。但在更复杂的图像数据(如MNIST的28x28像素,像素值0-255)上,标准化是必须的。
5.2 模型训练与基准测试
我们用自己实现的KNN和scikit-learn官方的KNN进行对比,这是一个很好的验证我们代码正确性的方法。
from sklearn.model_selection import train_test_split from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns from sklearn.neighbors import KNeighborsClassifier # 划分数据集 X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y) # 使用我们自实现的KNN my_knn = KNNClassifier(n_neighbors=5, weights='distance', metric='euclidean') my_knn.fit(X_train, y_train) y_pred_my = my_knn.predict(X_test) # 使用scikit-learn的KNN sk_knn = KNeighborsClassifier(n_neighbors=5, weights='distance', metric='euclidean') sk_knn.fit(X_train, y_train) y_pred_sk = sk_knn.predict(X_test) # 比较准确率 from sklearn.metrics import accuracy_score print(f"自实现KNN准确率: {accuracy_score(y_test, y_pred_my):.4f}") print(f"Scikit-learn KNN准确率: {accuracy_score(y_test, y_pred_sk):.4f}") # 输出详细的分类报告 print("\n自实现KNN分类报告:") print(classification_report(y_test, y_pred_my)) # 绘制混淆矩阵 cm = confusion_matrix(y_test, y_pred_my) plt.figure(figsize=(10,8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues') plt.xlabel('Predicted Label') plt.ylabel('True Label') plt.title('Confusion Matrix for Handwritten Digits (Our KNN)') plt.show()如果我们的实现正确,两个模型的准确率应该非常接近(可能因为随机数种子或细微实现差异有小数点后几位的差别)。混淆矩阵能帮助我们看清模型具体在哪些数字上容易混淆,比如“8”和“3”、“9”和“7”等。
5.3 特征工程与预处理的影响
在这个简单的数据集上,我们的模型可能已经表现不错。但我们可以尝试一些简单的预处理,看看能否提升性能。例如,我们可以尝试对像素值进行标准化。
from sklearn.preprocessing import StandardScaler scaler = StandardScaler() X_train_scaled = scaler.fit_transform(X_train) X_test_scaled = scaler.transform(X_test) # 注意:使用训练集的均值和方差来转换测试集 my_knn_scaled = KNNClassifier(n_neighbors=5, weights='distance') my_knn_scaled.fit(X_train_scaled, y_train) y_pred_scaled = my_knn_scaled.predict(X_test_scaled) print(f"标准化后的KNN准确率: {accuracy_score(y_test, y_pred_scaled):.4f}")对于这个特定的小数据集,标准化可能提升不大,甚至可能因为数据本身尺度统一而没有变化。但这个步骤在真实世界中至关重要。例如,如果你的数据包含“身高(米)”和“体重(公斤)”,不标准化的话,“体重”的微小变化在欧氏距离中产生的影响将远超“身高”,这显然不合理。
6. KNN算法的优势、局限与优化策略
6.1 算法优势与适用场景
KNN的优点非常突出,这也是它经久不衰的原因:
- 原理简单,直观易懂:不需要复杂的数学背景就能理解,非常适合教学和入门。
- 无需训练过程:
fit方法只是存储数据,训练时间复杂度为O(1)。对于需要频繁更新训练集(在线学习)的场景,KNN有天然优势,只需将新数据加入存储库即可。 - 对数据分布没有假设:不像线性回归要求线性关系,也不像朴素贝叶斯要求特征条件独立。它是一种非参数方法,能适应各种复杂的数据分布。
- 在多分类问题上表现自然:无需像一些二分类模型那样进行改造。
因此,KNN非常适合以下场景:
- 小规模数据集的快速原型验证:在项目初期,用KNN快速建立一个基线模型。
- 数据分布不规则或未知:当你不确定数据的内在结构时。
- 需要解释预测结果时:你可以直接展示“因为这几个样本和你最像,它们都是A类,所以我们也预测你是A类”,这种解释性在某些领域(如医疗辅助诊断)很有价值。
6.2 核心局限与性能瓶颈
KNN的缺点同样明显,在应用时必须心中有数:
- 计算复杂度高:预测时需计算与所有训练样本的距离,时间复杂度为O(n_samples_train * n_features)。当训练集很大(上百万)或特征维度很高(上千维)时,预测速度会慢到无法接受。
- 内存消耗大:需要存储整个训练集,对于大数据集不友好。
- 对高维数据效果差(“维数灾难”):在高维空间中,所有点之间的距离都趋于变得非常相似,导致“最近邻”的概念失去意义,模型性能急剧下降。
- 对不平衡数据敏感:如果某个类别的样本数量远多于其他类别,那么在进行多数表决时,新样本更容易被归为这个多数类,导致对少数类的预测精度极差。
- 对噪声和无关特征敏感:如果数据中存在大量噪声或与分类无关的特征,会严重影响距离计算,从而干扰邻居的选取。
6.3 常用优化与改进策略
针对上述问题,业界有一些常见的应对策略:
使用高效数据结构加速搜索:这是解决预测慢的核心方法。
- KD树:一种对k维空间中的点进行划分的二叉树结构,适用于低维空间(例如维度<20),可以將平均搜索复杂度从O(N)降低到O(log N)。
- 球树:KD树的改进,对高维数据更鲁棒。它将数据点组织成嵌套的超球体。
- 近似最近邻搜索:如Locality-Sensitive Hashing,通过哈希技术快速找到近似最近邻,用少量精度损失换取巨大的速度提升,适用于海量数据。scikit-learn的
KNeighborsClassifier默认在数据量大时会自动使用KDTree或BallTree。
特征选择与降维:用于应对“维数灾难”和无关特征。
- 使用过滤法(如方差选择、卡方检验)、包装法(如递归特征消除)或嵌入法(基于模型的特征重要性)选择最相关的特征子集。
- 使用主成分分析、线性判别分析或t-SNE等降维技术,将高维数据映射到低维空间,同时尽可能保留分类信息。
数据预处理:
- 标准化/归一化:处理量纲问题,前文已强调。
- 处理不平衡数据:对多数类进行欠采样,或对少数类进行过采样(如SMOTE算法),使类别分布更均衡。
调整距离度量与投票权重:
- 根据数据特性选择更合适的距离,如曼哈顿距离、余弦相似度、马氏距离等。
- 使用加权投票,让更近的邻居拥有更高权重,可以平滑噪声的影响。
实操心得:在实际项目中,我很少将KNN作为最终的生产模型,尤其是在数据量大或实时性要求高的场景。但它是我工具箱里不可或缺的“瑞士军刀”。我主要用它做两件事:一是项目初期的快速探索和基线建立,二是作为复杂模型(如集成模型)中的一个弱学习器。理解它的优缺点,能让你更清醒地知道何时该用它,何时该寻找更高级的算法。
7. 常见问题排查与调优技巧实录
在实际使用自实现或调优KNN时,你肯定会遇到各种各样的问题。下面我整理了一些典型问题及其排查思路,很多都是我自己踩过的坑。
7.1 预测结果全部为同一个类别
问题描述:无论输入什么数据,模型预测的标签都是同一个值(比如全是0)。
排查思路:
- 检查K值:首先确认你的K值是否设置得过大,比如K等于或超过了训练集中最少类别的样本数。如果K值过大,投票结果可能会被样本数最多的类别主导。
- 检查数据预处理:这是最常见的原因。确保你在预测前对数据进行了与训练时完全相同的预处理(如标准化)。一个典型的错误是:用原始数据训练,却把标准化后的数据拿去预测,或者相反。距离计算在完全不同的尺度上进行,结果必然失真。务必记住:测试集的标准化参数(均值和标准差)必须来自训练集,不能独立计算。
- 检查距离计算:在自实现代码中,仔细核对距离计算函数。例如,在计算欧氏距离平方时,是否错误地先开了方再求和?确保
np.sum的axis参数设置正确。 - 检查投票逻辑:特别是加权投票时,权重计算是否正确?是否存在除零错误导致权重为无穷大?打印出最近邻的标签和距离,手动验算一下投票过程。
7.2 模型准确率远低于预期
问题描述:在测试集或交叉验证中,准确率非常低,甚至低于随机猜测。
排查思路:
- 数据划分泄露:确保训练集和测试集是完全独立的。最常见的错误是在全局进行标准化(先对所有数据标准化,再划分训练测试),这会导致测试集信息“泄露”到训练过程中。正确的做法是:先划分,再分别用训练集的统计量去转换训练集和测试集。
- K值过小或过大:绘制“K值-验证集准确率”曲线,找到性能拐点。K=1时模型可能过拟合噪声,K过大则可能欠拟合。
- 特征尺度问题:再次强调,检查所有连续型特征是否经过了恰当的标准化/归一化。可以打印特征的最大最小值看看。
- 数据本身不可分:如果特征与标签之间几乎没有关联,任何模型都无能为力。检查特征与标签的相关性,或者用其他简单模型(如决策树)试试,如果大家都表现很差,那可能就是数据问题。
- 类别标签错误:检查训练数据的标签是否正确,是否存在标注错误。
7.3 预测速度异常缓慢
问题描述:模型预测一个样本需要好几秒甚至更久。
排查思路:
- 训练集规模:KNN的预测复杂度与训练集大小线性相关。如果训练集有几十万、上百万样本,预测慢是正常的。考虑是否可以使用子采样后的数据,或者必须转向使用KD树/球树等加速结构。
- 特征维度:高维特征会显著增加单次距离计算的开销。检查是否有大量无关或冗余特征,尝试进行特征选择或降维。
- 实现效率:在自实现代码中,确保距离计算使用了NumPy的向量化操作,避免低效的Python级循环。可以使用
%timeit魔法命令来 profiling 关键函数的耗时。 - 使用加速库:对于生产环境,考虑使用经过高度优化的库,如
scikit-learn(内部用Cython优化)、faiss(Facebook开源的相似性搜索库,针对大规模向量集做了极致优化)或annoy(Spotify开源的近似最近邻库)。
7.4 处理类别不平衡数据
问题描述:数据集中某些类别的样本数远多于其他类别,导致模型对少数类预测精度极差。
解决方案:
- 调整类别权重:在投票时,可以为不同类别的样本赋予不同的权重。例如,让少数类邻居的票数乘以一个大于1的系数。在scikit-learn中,
KNeighborsClassifier有一个class_weight参数可以设置为'balanced',它会自动根据类别频率调整权重。在我们自实现的加权投票中,可以手动融入这个逻辑。 - 重采样:
- 欠采样:随机删除一些多数类样本。风险是可能丢失重要信息。
- 过采样:复制少数类样本,或使用SMOTE等算法生成合成样本。风险是可能过拟合到少数类的噪声上。
- 改变决策阈值:KNN本身输出的是“票数”或“权重和”,你可以不直接采用“多数”原则,而是为少数类设定一个更低的获胜阈值。但这需要将KNN的预测过程修改为输出概率或置信度,实现起来更复杂。
一个简单的实践技巧是,在划分训练集时使用stratify参数(如train_test_split(..., stratify=y)),这可以保证训练集和测试集中的类别分布与原始数据集一致,至少能让你在评估时得到一个更可靠的基准。
