Jikipedia
第 41 篇

scikit-learn 之 kNN 分类

量化课堂第 41 篇(postId=3227,作者肖睿,编辑宏观经济算命师,难度进阶上、理解深度 level-1,2016-10-13 上线,v1.1 于 11-11 补漏字、v1.2 于 2016-11-30 改错字)。机器学习组第 6 篇,kNN 家族(38–41)的实操收官:把 38 篇的 kNN 原理落成 sklearn.neighbors.KNeighborsClassifier 的完整用法——每个参数是什么、fit 之后有哪些方法,再用「三团正态分布点」实例(参数正好复刻 38 篇的三种兔子)演示归一化→训练→预测→打分。阅读前提:kNN(38 篇,2227)、建议掌握 kd 树(40 篇,2843)。素材见 raw/collections/jq-quant-classroom/41-41-scikit-learn之knn分类md.md。

量化scikit-learnkNNKNeighborsClassifier归一化分类

这是什么

一篇「函数说明书」式的实操文,代码几乎全在正文。它把 KNeighborsClassifier 初始化的 8 个参数逐个解释,再讲 fit、kneighbors、predict、predict_proba、score 五个方法,最后给一个完整可跑的例子(600 个训练点 + 300 个测试点)演示怎么归一化、怎么打分,并画出分类图。诚实提醒贯穿全文:toy 数据的 97% 正确率在真实涨跌预测里达不到。

核心要点

KNeighborsClassifier 参数

  • 初始化签名(默认值):KNeighborsClassifier(n_neighbors=5, weights='uniform', algorithm='auto', leaf_size=30, p=2, metric='minkowski', metric_params=None, n_jobs=1)
  • n_neighbors = kNN 里的 k。
  • weights:'uniform' 等权投票;'distance' 按距离倒数加权;也可自定义。例子:最近的 3 点里有 1 个 A(很近)和 2 个 B(稍远)——等权 3NN 判 B;距离加权下若 A 的权重超过两个 B 之和则判 A。选哪个看应用场景。
  • algorithm:'brute'(蛮力)、'kd_tree'、'ball_tree'(另一类树结构)或 'auto'(自动选)。样本量与维度不同各有利弊,一般 'auto' 即可。
  • leaf_size:kd/ball 树的叶子大小;叶子里可以放多个点、到达叶子后在其中蛮算。多数场景影响不大,设 1 就行(教程默认给 30)。
  • metric + p:默认 'minkowski' + p=2,即 L₂(欧氏)距离;其他 metric 见官方文档。metric_params 是特殊 metric 需要的参数,默认 None。
  • n_jobs:并行线程数,默认 1,输入 −1 则用满 CPU 核。

fit 与常用方法

  • fit(X, y):X 是样本矩阵(每行一个样本,各样本长度须一致 = 特征数);y 是对应标签。fit 后按 algorithm 选项生成 kd/ball 树('brute' 则什么都不做)。
  • kneighbors(X=None, n_neighbors=None, return_distance=True):返回 (dist, index)。X 不传则默认用训练数据;index[i] 是训练数据里离 X[i] 第 j 近的元素下标,dist[i][j] 是对应距离。return_distance=False 时只返回 index。
  • predict(X):对每个样本返回预测的类别标签。
  • predict_proba(X):返回每个样本属于各类的概率 p[i][j]。类别顺序按词典排序:训练标签里若有 1、'1'、'a',则 1 是第 0 类、'1' 第 1 类、'a' 第 2 类(Python 中 1<'1'<'a')。
  • score(X, y):返回正确率。防过拟合的标准做法:把有标签样本分成训练/测试两组,一组学、一组测。

实际例子(复刻兔子的三团正态分布)

  • 生成 3 类各 200 个正态点:(50,5)、(30,4)、(45,2.5)(均值/方差对应 38 篇的悲伤/痛苦/绝望),labels=[1]×200+[2]×200+[3]×200。
  • 归一化(38 篇的坑):算 x_diff=max−min、y_diff,坐标除以各自 diff 再两两配对 → xy_normalized。
  • clf = neighbors.KNeighborsClassifier(30)clf.fit(xy_normalized, labels)
  • kneighbors:问 (50,5) 与 (30,3) 的最近 10 个 → 返回训练数据下标(如 97/134/177/144/10 与 278/569/242/324/504)。
  • predict:(50,5)→1 类,(30,3)→2 类,即 array([1, 2])。
  • predict_proba:(50,5)→[1,0,0](100% 类 1);(30,3)→[0,0.8,0.2](80% 类 2、20% 类 3)。
  • score:另生成 100×3 测试点 → 30NN 正确率 97%;换 1NN → 降到 95%(过拟合)。
  • 诚实提醒:精度高是因为训练/测试都是人为按正态分布生成的;真实场景(比如涨跌预测)很难达到这个精度。

画分类图(教学附加)

  • np.meshgrid 生成网格 → np.ravel() 拉直 + np.c_ 拼成坐标列 → predict → reshape 回网格 → plt.pcolormesh 上色;predict_proba 取单类概率画热力图。只能展示两个维度,教学用。

机制 / 论证

  • weights 为什么能改判:投票从「数人头」变成「按距离倒数的权重」——更近的样本说话更响,对边界/噪声更稳。A 近而 B 远时,等权会输给人数,距离加权能翻盘。
  • algorithm 三选一的本质:kNN 的 fit 几乎不「学」,只是存数据(brute)或建索引树(kd/ball);'auto' 让 sklearn 内部按数据规模/维度权衡。
  • 为什么归一化必须自己做:sklearn 的距离计算不感知量纲,38 篇的「轴比例失衡」在库里同样存在——所以本例先手动除以 x_diff/y_diff。
  • 1NN vs 30NN 打分差 = 偏差-方差的可测版:k 越大越平滑(30NN 97%),k=1 贴历史过拟合(95%),把 38 篇的定性说法量化了。

可操作

  • sklearn kNN 五步:
    1. 特征归一化(每轴除以 max−min 或等价缩放),否则距离被大量纲轴主导。
    2. clf = neighbors.KNeighborsClassifier(n_neighbors=k)(一般 algorithm 用默认 'auto')。
    3. clf.fit(X_train, y_train)
    4. clf.predict(X_test)clf.predict_proba(X_test)
    5. clf.score(X_test, y_test) 看测试正确率——训练/测试必须分开,防过拟合。
  • import 写法:from sklearn import neighbors;若只 import sklearn,则写 sklearn.neighbors.KNeighborsClassifier(...)
  • 归一化后预测时也要把新点除以同一个 x_diff/y_diff(本例 kneighbors/predict 的坐标都先归一)。
  • 涨跌预测等真实任务精度远低于 toy 的 97%——别拿人工数据的分数当真。
  • 本文是纯 sklearn 用法,无聚宽策略代码;配套 notebook(blog+kNN scikit-learn.ipynb,notebookCloneCount=321)在文末研究模块。

术语

  • KNeighborsClassifier:scikit-learn 的 k 最近邻分类器。
  • weights:近邻投票的权重方式(等权 / 距离倒数加权 / 自定义)。
  • algorithm:brute / kd_tree / ball_tree / auto。
  • leaf_size:树叶子容量(叶子内蛮算)。
  • kneighbors / predict / predict_proba / score:查询最近邻、预测、概率预测、正确率打分。

不确定 / 待验证

  • 示例代码把 random.normal(...) 写成 random,实际应是 np.random.normal(numpy)的概率写法笔误 [需要验证]。
  • 实例数字(97%/95%、kneighbors 下标、predict_proba [0,0.8,0.2])照录 raw;正文无随机种子,无法完全复算一致。
  • score 的测试集与训练集同分布生成,测的是「同分布内泛化」;跨市场/跨时段的外推性没讨论——呼应原文「真实涨跌预测达不到」。
  • 归一化只有「除以 max−min」一种示范,与 z-score 等标准化方法的关系没展开。
  • 画图部分(meshgrid/pcolormesh)只是教学演示,不是策略;>3 特征就画不出来。

相关

更新 2026-09-06

检索知识库

按标题、类型或正文检索