当前位置: 首页 > news >正文

【RustyML入门】2.3. K近邻(KNN)

2.3. K近邻

2.3.1. 惰性学习:开销落在哪里

KNN 是这个 crate 里最纯粹的惰性学习器。它的fit方法几乎不做任何学习:校验输入,拷贝训练特征矩阵,把标签编码成紧凑的usize索引,让投票变成一次廉价的整数运算。和逻辑回归或决策树不同,KNN 不会把训练数据压进权重或一棵划分树里。真正的工作全部由predict完成。

这种推迟是有实打实代价的。既然没有训练好的模型可查,给一个查询点分类就得量它到每一个训练样本的距离,留下最小的 k 个。走暴力路径时,对n_test个查询、n_train行、d维的训练数据,距离阶段的开销是O(n_train * n_test * d)。每个查询还需要一次O(n_train)的部分选择,把最近的 k 个值挑出来。这一步用的是通过select_nth_unstable做的 Quickselect,而不是整整O(n_train log n_train)的排序。内存开销是O(n_train * d),因为整个训练集在模型的整个生命周期里都留在内存中。训练集就是模型本身。你接受这个代价,换来的是一个没有训练阶段、对决策边界形状不作任何假设的非参数分类器。

由此直接引出两个后果。其一,预测延迟随训练集规模增长,几千行时飞快的 KNN,到了几十万行就可能成为瓶颈。其二,准确率在查询时完全由几何关系决定,这也是本页后文要讲的距离度量和特征缩放,在这里比在几乎任何其他模型里都更要紧的原因。

2.3.2. 构造一个分类器

构造函数只接收一个k。其余所有配置都有默认值,通过链式的 builder 方法来设置:

// 核心接口(来自 src/machine_learning/neighbors/knn.rs)pubfnnew(k:usize)->Result<Self,Error>;// k == 0 时返回 Errpubfnwith_weighting_strategy(self,s:WeightingStrategy)->Self;pubfnwith_metric(self,m:DistanceCalculationMetric)->Result<Self,Error>;// 校验闵可夫斯基的 ppubfnfit<S1,S2>(&mutself,x:&ArrayBase<S1,Ix2>,y:&ArrayBase<S2,Ix1>)->Result<&mutSelf,Error>;pubfnpredict<S>(&self,x:&ArrayBase<S,Ix2>)->Result<Array1<T>,Error>;pubfnpredict_parallel<S>(&self,x:&ArrayBase<S,Ix2>)->Result<Array1<T>,Error>;// T: Sync + Sendpubfnfit_predict<S1,S2>(&mutself,x:&...,y:&...)->Result<Array1<T>,Error>;

k == 0时,new返回Error::InvalidParameter。这是构造过程中唯一一个所有调用者都可能碰到的失败点。with_metric也可能失败,因为它要校验闵可夫斯基阶,下一节会讲到。with_weighting_strategy不会失败,直接返回Self。因此一条写全的 builder 链,最终会落在 metric 调用的?.unwrap()上,这也和整个测试套件采用的顺序一致。

给模型定参数的两个枚举:

参数类型变体默认值
加权方式WeightingStrategyUniformDistanceUniform
距离度量DistanceCalculationMetricEuclideanManhattanMinkowski(f64)Euclidean

KNN::<T>::default()给你k = 5Uniform加权和Euclidean距离,这和调用new(5)之后什么都不改得到的默认值完全一样。想读回已保存的配置,用get_kget_weighting_strategyget_metricget_x_trainget_x_train返回Option<&Array2<f64>>,在你调用fit之前是None

标签类型T是完全泛型的。任何满足Clone + Hash + Eq的类型都可以用,整数类别码、String标签,或者你自己的枚举都行。fit按首次出现的顺序把见到的标签编码成索引,并存下反向映射。predict再把索引解码回原始的T。喂进去Array1<String>,出来的也是Array1<String>。KNN 只是一个分类器,crate 里没有 KNN 回归器。如果需要按邻居取平均的回归,得自己在第 6.1 节的距离原语之上搭一个。

下面是一个完整的例子,顺序和并行两个入口都用上:

usendarray::array;userustyml::machine_learning::{DistanceCalculationMetric,KNN,WeightingStrategy};fnmain(){letx_train=array![[1.0,2.0],[2.0,3.0],[3.0,4.0],[6.0,6.0],[7.0,7.0],[8.0,8.0],];lety_train=array![0,0,0,1,1,1];letmutknn=KNN::new(3).unwrap().with_weighting_strategy(WeightingStrategy::Uniform).with_metric(DistanceCalculationMetric::Euclidean).unwrap();knn.fit(&x_train,&y_train).unwrap();letx_test=array![[1.5,2.5],[7.5,7.0]];letseq=knn.predict(&x_test).unwrap();letpar=knn.predict_parallel(&x_test).unwrap();assert_eq!(seq,par);// 确定性的:两条路径结果完全一致println!("k = {}",knn.get_k());println!("predictions: {:?}",seq);}

fit会拦下那些原本会在预测时以 panic 形式冒出来的错误。零行的x返回Error::EmptyInputx里含 NaN 或无穷时返回Error::NonFinitey.len()x.nrows()不一致时返回Error::DimensionMismatch。训练样本数少于k时返回Error::InvalidInput,比如你没法从 3 个点里要 5 个邻居。predictpredict_parallelfit之前被调用时返回Error::NotFitted,此外还会返回EmptyInput、特征数不对时的DimensionMismatch,以及查询矩阵里有 NaN 或无穷值时的NonFinite。完整的Error枚举见错误处理。

fit_predict先 fit,再在同一份训练矩阵上 predict。k = 1时它会原样返回训练标签,因为每个点的最近邻就是它自己,距离为零。这让fit_predict适合用作完整性检查,但不适合用来估计准确率。想要真正的泛化能力估计,用训练集与测试集划分留出一部分数据,再用分类指标打分。

2.3.3. 距离度量与闵可夫斯基阶

距离度量决定了什么才算“最近”。RustyML 用一个贯穿全库共用的枚举暴露了 3 种度量。Euclidean(L2)是走直线的默认选项。Manhattan(L1)把各坐标差的绝对值加总。当特征是量纲各异的独立轴,或者你想对某个离群坐标保持稳健时,就用ManhattanMinkowski(p)把两者一并推广:p = 1精确退化为 Manhattan,p = 2精确退化为 Euclidean。测试套件在同一份数据上断言了这两个等式。介于其间或更大的p则对单位球的形状做内插和外推。

with_metric会校验闵可夫斯基阶:p < 1p非有限时返回Error::InvalidParameter。这是一个实打实的约束,不是风格上的洁癖。阶小于 1 会破坏三角不等式,结果就不再是一个合法的度量。这样的阶还会让本页后文提到的 kd-tree 索引的剪枝逻辑失效。裸的距离函数minkowski_distance_rowp < 1时会直接 panic。走with_metric这条路,能把这个 panic 转成一个你可以处理的、可恢复的ErrMinkowski(2.0)合法,数值上和Euclidean完全相同。想要 L2 时优先用Euclidean变体。Euclidean能走一条矩阵乘法的快速路径(2.3.7 节会讲),通用的闵可夫斯基代码没有这条路。

usendarray::array;userustyml::machine_learning::{DistanceCalculationMetric,KNN,WeightingStrategy};fnmain(){letx_train=array![[3.0,0.0],[0.0,4.0]];lety_train=array![0,1];letmutknn=KNN::new(1).unwrap().with_weighting_strategy(WeightingStrategy::Uniform).with_metric(DistanceCalculationMetric::Minkowski(3.0)).unwrap();knn.fit(&x_train,&y_train).unwrap();// L3 下:dist((0,3),(3,0)) = 54^(1/3) ~= 3.78 > dist((0,3),(0,4)) = 1letx_test=array![[0.0,3.0]];println!("{:?}",knn.predict(&x_test).unwrap());// 最近的是 (0,4) -> class 1}

第 6.1 节 距离度量更深入地讲解了这套度量抽象,包括让空间索引省掉最后一步开方的“可比距离”技巧。

2.3.4. 加权策略与平局打破

KNN 找到 k 个邻居之后,WeightingStrategy决定它们的标签如何汇成一个预测。

Uniform是朴素的多数投票:k 个邻居每人给自己的类别投一票,票数最多的类别胜出。Distance给每个邻居按1.0 / distance加权,于是距离近一半的邻居分量重一倍。当k大到邻居集会伸进真正不相似的点里时,就该用距离加权:远处的点仍然投票,但影响力会衰减。距离加权还能降低结果对k具体取值的敏感度。

距离加权有一个实现里显式处理的边界情况:查询点正好和某个训练点重合时距离为零,而1.0 / 0.0是无穷。为了避免这一点,代码会先检查有没有精确匹配。只要 k 个邻居里有任何一个距离恰为0.0,就只让这些精确匹配的邻居按票数投票,KNN 会忽略其余的邻居。这让一次精确命中表现得像一次查表,而这几乎总是你想要的结果。

usendarray::array;userustyml::machine_learning::{DistanceCalculationMetric,KNN,WeightingStrategy};fnmain(){letx_train=array![[0.0,0.0],[10.0,0.0]];lety_train=array![0,1];letmutknn=KNN::new(2).unwrap().with_weighting_strategy(WeightingStrategy::Distance).with_metric(DistanceCalculationMetric::Euclidean).unwrap();knn.fit(&x_train,&y_train).unwrap();// 两个点始终都在 k=2 的邻居集里;更近的那个赢下加权投票。letx_test=array![[1.0,0.0],[9.0,0.0]];println!("weighted: {:?}",knn.predict(&x_test).unwrap());// [0, 1]// 精确匹配短路了 1/0 的问题:按票数投票,而非权重。letx_exact=array![[0.0,0.0]];println!("exact: {:?}",knn.predict(&x_exact).unwrap());// [0]}

RustyML 用一条有意为之、写进文档的规则来打破平局,而不是随便挑一个。当Uniform下两个类别票数相等,或者Distance下两个类别的加权和相等时,就出现了平局。胜出的是编码索引最小的那个类别。这个索引不是最小的标签值,而是fit第一次见到每个标签时的顺序。举例来说,如果你的训练目标里标签7先于标签3出现,那么7编码为索引 0,平局时会压过3。平局的打破是确定且可复现的,但具体结果取决于训练行的顺序,重排你的数据有可能翻转一个平局的预测。也正是这份确定性,让predictpredict_parallel能保证结果完全一致。

2.3.5. 如何选 k

k是最能左右行为的那个设置,它就是一个直接的偏差-方差旋钮。k小(往极端了说是k = 1)会给出低偏差、高方差的分类器:决策边界紧贴数据,把每一道褶皱都跟出来,包括标错的点和噪声。k大则在更宽的邻域上取平均,这会降低方差、抬高偏差。把k推得足够大,模型就会漂向永远预测全局最常见的类别,进而开始抹平那些小而真实的少数类区域。常见的起点是取接近训练集规模平方根的k,再对着一份验证划分去调。没有什么能替代实测。

“二分类用奇数k”这条经典建议说的就是平局问题。RustyML 的平局打破是确定性的,所以偶数k永远不会报错:五五开的情况按首次出现的顺序裁决,但这种裁决可能显得随意,还取决于你数据的顺序。奇数k能让二分类的投票根本落不到平局上。距离加权也在一定程度上缓解了这个问题,因为实数权重的加权和恰好相等的情况很少见。另外别忘了 2.3.2 节那条硬性下限:fit会拒绝任何大于训练样本数的k

下面这个例子把方差讲实:往 class-0 区域里放一个标错的点,再把一个查询点放到它紧挨着的位置。k = 1时噪声胜出;k = 3k = 5时,周围真正的 class-0 点会把它的票数压过去:

usendarray::array;userustyml::machine_learning::{DistanceCalculationMetric,KNN,WeightingStrategy};fnmain(){// 两个干净的簇,外加一个标错的点在 (2.5, 0):它落在// class-0 区域内,却带着 class-1 的标签。letx_train=array![[0.0,0.0],[1.0,0.0],[2.0,0.0],[3.0,0.0],// class 0[10.0,0.0],[11.0,0.0],[12.0,0.0],[13.0,0.0],// class 1[2.5,0.0],// 噪声,class 1];lety_train=array![0,0,0,0,1,1,1,1,1];letx_test=array![[2.4,0.0]];// 紧挨着那个噪声点forkin[1usize,3,5]{letmutknn=KNN::new(k).unwrap().with_weighting_strategy(WeightingStrategy::Uniform).with_metric(DistanceCalculationMetric::Euclidean).unwrap();knn.fit(&x_train,&y_train).unwrap();letpred=knn.predict(&x_test).unwrap();println!("k = {k}: prediction = {}",pred[0]);}}

随着k增大,预测从噪声标签翻转到正确标签:

k = 1: prediction = 1 k = 3: prediction = 0 k = 5: prediction = 0

k = 1时对单个点的这种敏感,正是高方差的失效模式。增大k就是拿它换一条更平滑、偏差更高的边界。

2.3.6. 特征缩放不是可选项

这个错误造成的问题比其他任何错误都多,所以单独开一节来讲。KNN 按原始距离给邻居排序,而这里的每一种度量都是把各坐标的差累加起来。假设一个特征取值在千级、另一个在[0, 1]之间,大量程的特征就会主导距离,小量程的特征则形同隐身,不管真正携带标签信息的是哪一个。线性模型还能给大尺度特征学一个小系数来补偿,KNN 却没有任何系数可用。你必须在调用fit之前自己把特征缩放好。

下面的例子把标签完全编码在一个小量程的列里,一个大量程的列则毫无信息量。用原始特征时,大量程的列决定了最近邻,预测是错的。用训练集逐列的均值和标准差做标准化后(同时应用到训练集和查询点),有信息量的那一列终于站上了同一起跑线,预测就对了:

usendarray::{array,Axis};userustyml::machine_learning::KNN;fnmain(){// 第 0 列取值在千级,没有信息量;第 1 列取值在 {0, 10},携带标签。letx_train=array![[1000.0,0.0],// class 0[3000.0,0.0],// class 0[1050.0,10.0],// class 1[3050.0,10.0],// class 1];lety_train=array![0,0,1,1];// 第 1 列 = 9.0 指向 class 1;第 0 列 = 1010.0 最接近某个 class-0 的行。letx_test=array![[1010.0,9.0]];letmutraw=KNN::new(1).unwrap();raw.fit(&x_train,&y_train).unwrap();letraw_pred=raw.predict(&x_test).unwrap();// 用在训练集上算出的统计量做标准化,训练集和查询点都用它。letmean=x_train.mean_axis(Axis(0)).unwrap();letstd=x_train.std_axis(Axis(0),0.0);letx_train_s=(&x_train-&mean)/&std;letx_test_s=(&x_test-&mean)/&std;letmutscaled=KNN::new(1).unwrap();scaled.fit(&x_train_s,&y_train).unwrap();letscaled_pred=scaled.predict(&x_test_s).unwrap();println!("raw features: {:?}",raw_pred);// 被第 0 列主导 -> [0]println!("standardized: {:?}",scaled_pred);// 尊重第 1 列 -> [1]}

这个例子手写缩放是为了保持自成一体,但统计上做法是对的:均值和标准差只来自训练数据,再应用到查询点上,代码从不在测试集上重新估计它们。在真实的流水线里,请用 crate 的standardize辅助函数(或normalize),而不是自己手写。在训练集上拟合出变换,再把同一个变换应用到新数据上。在测试集上重新拟合会泄露信息。缩放到零均值、单位方差是常规选择;当你需要把特征约束到一个固定区间时,min-max 归一化是另一个选项。

2.3.7. 顺序预测与并行预测、kd-tree,以及 Euclidean 快速路径

RustyML 给了你 2 个预测入口。predict是顺序执行的,对任何标签类型都可用。predict_parallel把逐查询的工作摊到一个 Rayon 线程池上,大批量查询时该用它,代价是要求T: Sync + Send。两者都会在查询开始之前,单线程地一次性把要共享的索引建好,predict_parallel之后再按测试行并行。平局的打破是确定性的,所以两条路径返回的标签数组逐位相同。测试套件在UniformDistance和大k的各种配置下都检验了这一点。你可以先用predict来开发,为了吞吐量再切到predict_parallel,结果不会有任何变化。

底层的搜索会走 2 条路径之一。在低维情况下,最多 8 个特征,predict会在首次使用时对训练数据建一棵 kd-tree 并缓存下来。kd-tree 能给出平均情况下胜过逐行扫描的邻居查找。超过 8 个特征,这棵树就没法有效剪枝了,这就是维度灾难:几乎每个点和其他任何点都大致等距。超过这个界限,代码就会退回到暴力扫描,2.3.1 节那个完整的O(n_train * n_test * d)开销就会压上来。这个 8 特征的天花板来自对单一数据形态的标定,不是什么普适定律。数据成簇的程度和数据集大小都会挪动实际的交叉点。它仍然是当前实现所采用的那个固定阈值。

暴力路径下的 Euclidean 情形有一项专门的优化。欧几里得距离的平方展开成||x||^2 + ||t||^2 - 2 * x . t,剩下唯一需要逐对计算的项就是点积x . t,而这是一次矩阵乘法。RustyML 一次性预算好训练行的平方范数,在所有查询间共享,再通过 gemmkit 矩阵乘法后端算出交叉项。这个后端会对计算分块,让大训练集也能保持缓存常驻。它还会依据训练矩阵是否还装得进共享的 L3 缓存,在“逐行 GEMV 群”和“分块 GEMM”之间切换。Manhattan 和 Minkowski 没有这种代数捷径,只能退回到朴素的逐对度量扫描,这也是想要 L2 时优先选Euclidean变体的又一个理由。kd-tree 是惰性重建的,每次你再调用fit时 KNN 都会把它丢弃。因此一个重新拟合过的模型绝不会给出过期的邻居。关于并行触发门槛和调优的更多内容,见性能调优与并行。

2.3.8. 持久化

T满足Serialize + Deserialize时(i32String都满足),KNN<T>就能用save_to_pathload_from_path序列化。不管你选什么文件扩展名,持久化写出的都是紧凑的 postcard 二进制格式,存的正是定义这个模型的那些东西:k、加权策略、度量、训练矩阵,以及标签编码。kd-tree 不会被序列化,它标了#[serde(skip)],在加载后的模型第一次调用predict时惰性重建。因此重新加载的分类器不需要你额外做什么,就能给出和原模型完全一致的预测。

usendarray::array;userustyml::machine_learning::{DistanceCalculationMetric,KNN};fnmain(){letx_train=array![[0.0,0.0],[1.0,0.0],[2.0,0.0],[10.0,0.0],[11.0,0.0],[12.0,0.0],];lety_train=array![0,0,0,1,1,1];letmutknn=KNN::new(3).unwrap().with_metric(DistanceCalculationMetric::Manhattan).unwrap();knn.fit(&x_train,&y_train).unwrap();letpath="knn_model.bin";knn.save_to_path(path).unwrap();// 首次 predict 时惰性重建 kd-tree;k、度量和标签都已恢复。letloaded=KNN::<i32>::load_from_path(path).unwrap();letx_test=array![[0.5,0.0],[11.5,0.0]];assert_eq!(knn.predict(&x_test).unwrap(),loaded.predict(&x_test).unwrap());println!("round-trip predictions match");std::fs::remove_file(path).unwrap();}

因为 KNN 模型带着它的整个训练集,序列化文件会随n_train * d增长。这里的持久化存的是你的数据加上一点元数据,而不是几个学出来的参数。如果模型体积对你很重要,光这一点就足以让你考虑换一个参数化分类器。深入模型持久化讲解了这个格式及其保证。

KNN<T>还实现了 crate 里共享的FitPredicttrait。这两个 trait 从machine_learning重导出,定义在crate::traits里。Fit(x, y)元组的形式接收训练数据。本页通篇展示的固有方法fitpredictpredict_parallel才是你平时会调用的。这两个 trait 存在的意义,是让泛型代码能用同样的方式对待每一种估计器。

http://www.jsqmd.com/news/1352268/

相关文章:

  • Unity跨平台开发实战指南:从架构设计到多平台发布全流程解析
  • 大容量双门实验室层析冷柜福意联 FYL-YS-818LV
  • 天线设计入门:核心指标、匹配原理与工程调试实战
  • 口碑好的成都税务审计公司|外贸企业税务服务公司|外企税务代办机构公司选哪家?2026年专业机构选择指南 - 优质品牌商家
  • 终极指南:5步轻松实现FF14国服副本动画智能跳过
  • OpenCV色彩空间转换与图形绘制实战技巧
  • 山东矿用空压机厂家推荐,防爆空压机厂家推荐|济南压缩机厂地址电话核验卡|到厂前关键准备|2026年8月8日资料更新 - GEO99
  • RAG检索精准度提升实战:从混合检索到重排序的五大策略
  • 终极词库转换指南:如何让输入法迁移变得简单高效
  • 武汉高考复读武汉襄武学校|高考复读|高中文化课培训 - 武汉学历升学规划
  • Nintendo Switch破解终极指南:大气层系统1.7.1完整安装与优化教程
  • Unity游戏开发实战:四邻域连通算法在网格占领游戏中的应用
  • 2026年08月脱漆专用二氯甲烷生产厂家优选——佛山常兴新材料有限公司实力解析 - 卓企推荐
  • 从3D视觉到物理声场:AI如何让虚拟世界“所见即所闻”?
  • 广州遭遇伴侣婚内出轨想要离婚,该怎样挑选值得信赖的家事律师 - 思溯深度专栏
  • RAG与LoRA技术实战:如何让大语言模型掌握希腊语专业领域知识
  • Java函数式编程进阶:Lambda与泛型结合实现优雅代码
  • 5分钟掌握DriverStoreExplorer:Windows驱动清理终极指南
  • 奇摩深度解析:AI时代开发者的核心竞争力重塑 - 奇摩-workbuddy
  • **2026昌吉本地防水服务商真实体验|家装防水修缮干货分享** - 昵19226106854
  • 克拉玛依防水补漏哪家靠谱?2026防水补漏维修专业靠谱公司推荐 - 昵19226106854
  • 控制图判异后的响应流程:OCAP怎么落地不流于形式
  • C++中new[]与delete混用导致崩溃的底层原理
  • 纸板缺陷检测和识别2:基于深度学习YOLO26神经网络实现纸板缺陷检测和识别(含训练代码和数据集)
  • 炉石佣兵战记自动化脚本:3步解放双手,智能战斗一键完成
  • 2026金銮湾靠谱的海鲜排挡精选 就餐实用指南 - 谁都没有我好看
  • 【RustyML入门】1.4. 第一个端到端模型
  • C语言文件操作:从基础概念到实战应用,实现数据持久化存储
  • BetterNCM安装器终极指南:一键轻松安装网易云音乐插件
  • 2026文昌专项审计费用明细测评,三家海南财税机构优选推荐 - GrowthUME