2.9. MeanShift
MeanShift 是一种聚类算法。当你不知道数据里有几个簇、又不想去猜一个数字时,就用它。和 KMeans 不同,KMeans 需要你事先给出 k,MeanShift 则从数据自身的密度中找出簇的数量。它只需要你提供一个输入:带宽(bandwidth)。算法让每个种子点沿着密度曲面向上移动,直到停在某个模态(mode)上。
存活下来的模态就是簇中心,它们的数量由数据决定,而不是由你指定。这份自由是有代价的:这个代价就是带宽。它是唯一一个控制全局的设置,把它调错,也是得到糟糕结果的唯一途径。
RustyML 通过 MeanShift 以及一个独立的函数 estimate_bandwidth 提供这个算法,两者都从 rustyml::machine_learning 重新导出。核函数、合并规则、噪声标签,以及带宽估计器,全都与 scikit-learn 1.9.0 一致;在一个 10 点的参考数据集上,簇中心、它们的编号,以及每一个标签,都与 scikit-learn 的输出完全相同。
2.9.1. 用平坦核寻找模态
想象在你的点集上铺了一层核密度估计:每个样本都贡献一个小鼓包,样本聚集的地方,这些鼓包便叠加成峰。
MeanShift 把一个候选中心放在某个种子位置上,然后用周围数据的均值替换它。这个均值所在的方向,正是密度上升的方向。反复迭代之后,中心便爬到最近的峰,也就是一个模态(mode)。从许多种子出发各做一遍,就能找出所有模态。
用哪种核给这个均值加权,决定了算法的具体行为。RustyML 用的是平坦核 (flat kernel),和 scikit-learn 一样:对中心 c,落在它 bandwidth 之内的点各计一次,之外的点计零。于是下一个中心就是球内那些点的朴素均值:mean{ x_i : ||c - x_i|| <= bandwidth }。
这个窗口是一个半径为 bandwidth 的硬球,收敛时球里装了多少个点,就是这个模态的强度(intensity)。下面的合并阶段正是按这个强度给模态排序的。
本 crate 早先的版本用的是在整个数据集上加权的高斯核。RustyML 移除 了这个高斯核,而不是把它保留为一个选项:它在 scikit-learn 里没有对应物,因而没有任何东西能验证它,而且它也给不出合并规则所需的窗口点数。MeanShift 只有一种核,也没有 kernel 参数。
当带宽非常小时,平移后的中心,其球内可能变空。一种粗暴的做法是照常除以零,把中心塌缩到原点,这样会在 (0, 0, ...) 处凭空塞进一个和任何真实点都无关的假簇。RustyML 不这么做,而是把中心留在原地并停止迭代,这与 scikit-learn 在邻域为空时采取的提前退出完全一致。
一个没有邻居的点,会自成一个模态。正因如此,把带宽一路收缩到接近零时,算法会优雅地退化成"每个孤立点自成一簇",而不是给出没有意义的输出。
等所有种子都收敛之后,MeanShift 用 scikit-learn 的方法去除重复的模态:按强度排序的贪心抑制。具体做法是:先按各自窗口里装了多少个点,从多到少给这些收敛后的模态排序;然后沿着这个顺序走,保留一个模态,并丢弃与它相距不到一个 bandwidth 的其他所有模态。
这样保留下来的中心,才是一个真正的密度模态。早先的版本会把被抑制的模态平均进保留的中心里,这种做法会把中心从密度峰值上拽开,还会让结果取决于这一遍处理种子的顺序。
存活下来的中心数量,就是簇的数量,这个数字完全由数据和带宽决定,你从不需要直接指定它。每个输入样本,都会被打上离它最近的存活中心的标签。
正是这种两阶段结构------先让许多种子各自收敛,再合并它们------解释了为什么稍微偏大的带宽仍然往往能给出干净的结果:即便种子收敛到略有出入的位置,合并这一步也会把它们并到一起。
2.9.2. 构造 MeanShift
MeanShift::new 接收带宽并返回一个 Result:一个非正、非有限的带宽是使用错误,RustyML 不会悄悄把它夹到合法范围里:
rust,ignore
let ms = MeanShift::new(2.0)? // 唯一的必填参数
.with_max_iter(300)? // 返回 Result------校验 > 0
.with_tolerance(1e-4)? // 返回 Result------校验为正且有限
.with_bin_seeding(true) // 返回 Self------不会失败的开关
.with_cluster_all(true); // 返回 Self------不会失败的开关
返回类型上的这种区分是有意为之的,也很容易被忽略:两个收敛相关的 setter 会校验参数,交还 Result<Self, Error>,因此要配合 ? 使用;两个布尔开关不会失败,交还的是 Self,可以直接链式调用。
MeanShift::default() 等价于 new(1.0),其余参数都取默认值。这在初次尝试时很方便,但它几乎从来不是你真正想要的带宽,详见 2.9.4。
| 参数 | 构造 / 设置方法 | 默认值 | 含义 |
|---|---|---|---|
bandwidth |
new(bandwidth) |
无(必填) | 平坦核窗口的半径,同时也是合并半径和离群点的判定阈值。必须为正且有限。 |
max_iter |
with_max_iter |
300 |
每个种子的迭代上限;当某个种子始终达不到 tol 时,为最坏情形封顶。必须非零。 |
tol |
with_tolerance |
1e-3 |
收敛阈值;一旦某个种子的平移长度小于它,该种子即停止。必须为正且有限。 |
bin_seeding |
with_bin_seeding |
false |
通过把空间分箱到网格上来缩减种子集(见 [2.9.5](#参数 构造 / 设置方法 默认值 含义 bandwidth new(bandwidth) 无(必填) 平坦核窗口的半径,同时也是合并半径和离群点的判定阈值。必须为正且有限。 max_iter with_max_iter 300 每个种子的迭代上限;当某个种子始终达不到 tol 时,为最坏情形封顶。必须非零。 tol with_tolerance 1e-3 收敛阈值;一旦某个种子的平移长度小于它,该种子即停止。必须为正且有限。 bin_seeding with_bin_seeding false 通过把空间分箱到网格上来缩减种子集(见 2.9.5)。 cluster_all with_cluster_all true 把每个点都分配到某个簇;设为 false 则把过远的点标为 -1。))。 |
cluster_all |
with_cluster_all |
true |
把每个点都分配到某个簇;设为 false 则把过远的点标为 -1。 |
非法参数会从 new、with_max_iter、with_tolerance 以 Error::InvalidParameter 的形式返回,正是 1.6. 错误处理 里讲过的那个错误类型。
2.9.3. 拟合、预测与读取结果
fit 接收一个二维数组,每行一个样本,跑完算法后返回 &mut Self;predict 把新的点映射到学到的中心上,返回 Array1<isize>;fit_predict 一次做完这两件事,直接把训练集的标签返回给你。
标签之所以是有符号的,是因为 -1 被留给了噪声,这与 DBSCAN 和 scikit-learn 的约定一致。正因如此,任何聚类估计器的输出都能不加转换地喂给 5.3. 聚类指标 里的任何一个指标。fit 之后,拟合过程发现的一切都能通过各个 getter 取出。
rust
use ndarray::Array2;
use rustyml::machine_learning::MeanShift;
fn main() {
// 两个紧凑的点团:5 个点在 (0, 0) 附近,5 个点在 (20, 20) 附近。
let data = Array2::from_shape_vec(
(10, 2),
vec![
-0.1, 0.0, 0.1, 0.0, 0.0, -0.1, 0.0, 0.1, 0.0, 0.0, // 点团 A
19.9, 20.0, 20.1, 20.0, 20.0, 19.9, 20.0, 20.1, 20.0, 20.0, // 点团 B
],
)
.unwrap();
let mut ms = MeanShift::new(2.0).unwrap();
let labels = ms.fit_predict(&data).unwrap();
let centers = ms.get_cluster_centers().unwrap();
println!("clusters found: {}", centers.nrows()); // 由数据自行得出:2
println!("labels: {:?}", labels);
println!("samples per center: {:?}", ms.get_n_samples_per_center().unwrap());
println!("iterations run: {}", ms.get_actual_iterations().unwrap());
}
这些 getter 分成两类:一类给结果,一类回显配置。get_cluster_centers 返回 Option<&Array2<f64>>,每个簇一行;get_labels 返回 Option<&Array1<isize>>;get_n_samples_per_center 返回 Option<&Array1<usize>>,给出分配到每个中心的输入样本数。当 cluster_all = true 时,这些计数之和等于样本总数;当 cluster_all = false 时,被标为 -1 的离群点不计入其中。
get_actual_iterations 返回 Option<usize>,是所有种子里最大的迭代次数,据此你能判断这次运行是收敛了还是撞上了 max_iter。这 4 个 getter 在 fit 之前都是 None,之后才是 Some。
其余的 getter------get_bandwidth、get_max_iterations、get_tolerance、get_bin_seeding、get_cluster_all------只是把配置读回来而已。
predict 有 4 种失败情形:在 fit 之前调用它会返回 Error::NotFitted;传入空数组会返回 Error::EmptyInput;传入特征数与训练数据不一致的点会返回 Error::DimensionMismatch;传入包含 NaN 或无穷值的数据会返回 Error::NonFinite。这几种都不是靠重试能挽回的,把它们当成在运行时暴露出来的编程错误来对待。
2.9.4. 带宽:决定一切的超参数
一次 MeanShift 运行的方方面面都源自带宽:它是平坦核窗口的半径,也是合并阶段的半径,还是 cluster_all 关闭时判定离群点的阈值。带宽控制着每个种子能看多远、邻近模态被多强烈地并成一个,以及噪声从哪里开始。
带宽太小会导致过度切分:种子够不到一个簇自身的分布跨度,模态越冒越多,极端情况下每个彼此分离的点都自成一簇。带宽太大则会导致切分不足:相距很远的点团互相拉扯,直到它们的模态漂到一处,合并这一步把它们熔成一个,最终整个数据集变成了一个簇。
没有哪个默认值对任意数据都合适,因为正确的取值是一个长度尺度,量纲和你的特征相同。
手头没有任何先验估计时,estimate_bandwidth 能给出一个由数据驱动的起点。它接收数据和一个可选的 quantile(默认 0.3),还接收一个可选的子采样规模 n_samples(默认取全部行,并夹到数据集大小)以及一个可选的 random_state。
它计算 k = max(1, floor(n * quantile)),然后度量每个点到自己第 (k - 1) 近邻的距离,返回这些距离的均值 。这是一个局部密度统计量,回答的是"一个典型的点离自己邻域的边缘有多远",而这正是带宽该有的含义。
(k - 1) 这一项复现了 scikit-learn 的差一:它的近邻查询把查询点自身也数了进去。结果与 scikit-learn 1.9.0 吻合到 1e-14 以内,邻域里只有一个点时返回 0.0,和 scikit-learn 一样。
早先的版本返回的是全体两两距离分布的一个分位数。那是一个全局离散度指标,在有簇结构的数据上远大于带宽该有的量级,会把一切都并成一个簇。如果你曾对着旧的估计器调过带宽,现在请重新估计。
0.3 的分位数给出的是一段典型的中短邻域半径,通常落在簇内尺度附近。分位数越大,越偏向更大的带宽和更少的簇。
rust
use ndarray::Array2;
use rustyml::machine_learning::{MeanShift, estimate_bandwidth};
fn main() {
// 三个彼此分离的点团,每个 12 个紧凑的点。
let mut v: Vec<f64> = Vec::new();
for (cx, cy) in [(0.0, 0.0), (10.0, 0.0), (5.0, 9.0)] {
for k in 0..12u32 {
v.push(cx + ((k * 7) % 5) as f64 * 0.05 - 0.1);
v.push(cy + ((k * 3) % 5) as f64 * 0.05 - 0.1);
}
}
let data = Array2::from_shape_vec((36, 2), v).unwrap();
// 一个直接来自数据的合理起点。
let bw = estimate_bandwidth(&data, Some(0.3), None, Some(0)).unwrap();
println!("estimated bandwidth: {:.3}", bw);
// 扫一遍:小带宽过度切分,大带宽把一切并成一个。
for bandwidth in [0.05_f64, 0.5, 3.0, 30.0] {
let mut ms = MeanShift::new(bandwidth).unwrap();
ms.fit(&data).unwrap();
let k = ms.get_cluster_centers().unwrap().nrows();
println!("bandwidth {bandwidth:>5} -> {k} clusters");
}
}
簇的数量随带宽变化的方向和你预期的一致。由于确切的数量取决于数据,下面的输出请当成大致形态来看,而非字面数字:
text
estimated bandwidth: <small positive value>
bandwidth 0.05 -> many clusters (blobs fragment; over-segmentation)
bandwidth 0.5 -> one cluster per blob
bandwidth 3 -> one cluster per blob
bandwidth 30 -> a single cluster (all blobs merged; under-segmentation)
实操流程很简单:调一次 estimate_bandwidth,用它给的值去拟合,看看簇的数量;簇太多就调高带宽,太少就调低。用一个不需要真实标签的指标来验证这个选择,比如 5.3. 聚类指标 里的轮廓系数。
因为 estimate_bandwidth 直接度量距离,当各特征的量纲不一致时,先对特征做标准化(见 4.2. 标准化与归一化)能让这个估计更有意义。
2.9.5. Bin seeding、cluster_all 与离群点标签
默认情况下,每个输入点都是一个种子,这是最彻底的做法。正如 2.9.7 所说,这也是拟合具有确定性的原因。但在稠密数据集上,用每个点当种子很浪费,因为一个点团里成千上万个种子最终都会爬向同一个模态。
with_bin_seeding(true) 解决了这个问题:它把特征空间量化到一个网格上,格子的边长为 bandwidth,然后每个非空格子只保留 1 个代表性种子。种子少了,往上爬的次数就少,拟合也更快。
代价是种子布置更粗糙、更近似:某个模态的吸引域里如果始终没有一个网格代表点,它就可能被漏掉。当种子循环占据了大部分运行时间、且数据稠密到整格整格地被填满时,bin seeding 才是值得做的取舍;在小规模或稀疏的数据上,它省不下多少,反而只会牺牲分辨率。
cluster_all 决定那些并不真正属于任何模态的点该如何处置。取默认的 true 时,MeanShift 会把每个点都强行归到最近的中心,标签总落在 0..n_clusters 里,没有噪声这个概念。把它设为 false,任何离所有中心都超过一个 bandwidth 的点,都会被打上 -1 ------这是 scikit-learn 的噪声取值,也正是 DBSCAN 用的那个。这条规则对 fit 得到的训练标签和 predict 给出的新点标签同样适用。
-1 取代了早先那个等于 n_clusters 的哨兵值:对任何在下游数不同标签个数的代码来说,那个哨兵会被当成一个真实的额外簇。如果你的代码在拿标签和簇数量作比较,请改成 label < 0。
rust
use ndarray::Array2;
use rustyml::machine_learning::MeanShift;
fn main() {
let data = Array2::from_shape_vec(
(10, 2),
vec![
-0.1, 0.0, 0.1, 0.0, 0.0, -0.1, 0.0, 0.1, 0.0, 0.0,
19.9, 20.0, 20.1, 20.0, 20.0, 19.9, 20.0, 20.1, 20.0, 20.0,
],
)
.unwrap();
let mut ms = MeanShift::new(2.0).unwrap().with_cluster_all(false);
ms.fit(&data).unwrap();
// (10, 10) 离两个点团都约 14 个单位,远超 2.0 的带宽。
let probe = Array2::from_shape_vec((1, 2), vec![10.0, 10.0]).unwrap();
let pred = ms.predict(&probe).unwrap();
if pred[0] < 0 {
println!("outlier: label {}", pred[0]); // -1
} else {
println!("assigned to cluster {}", pred[0]);
}
}
启用 cluster_all = false 时,请判断 label < 0,不要想当然地以为标签是连续密集的。万一所有点都成了噪声,朴素的 labels.iter().max() 就不再能告诉你簇的数量了,可靠的计数是 get_cluster_centers().unwrap().nrows()。
2.9.6. 开销、收敛与并行
MeanShift 是平方复杂度的,把它用到大数据集之前,先把这一点想清楚。每个种子的每次迭代都要在 d 维里触碰全部 n 个点,判断哪些落在窗口内并对它们求均值,所以单个种子的开销是 O(iterations * n * d)。用默认的播种方式------每个点都是种子------一共有 n 个种子,一次完整拟合大致就是 O(iterations * n^2 * d)。
这和 DBSCAN 的两两扫描属于同一渐近量级,比 KMeans 的 O(iterations * n * k * d) 重得多,因为 KMeans 的 k 通常远小于 n。bin seeding 把种子数从 n 削减到非空格子的数量,这降低了常数项,但改变不了每次迭代那个平方项。
收敛的上界是 max_iter(默认 300):某个种子一旦平移量降到 tol 以下就会提前停下。get_actual_iterations 报告的是所有种子里需要的最大迭代次数,因此这个值卡在 max_iter 上,就是在提示你有种子始终没能安定下来。
实现只在划算的时候才并行。在 fit 里,各个种子的向上攀爬彼此独立,一旦总工作量------种子数乘样本数乘特征数------越过 RustyML 校准好的扫描级门槛,这些攀爬就会铺到 Rayon 线程池上跑。这个门槛默认是 262,144 次元素运算,可以通过 crate::tuning 调整。低于这个门槛,分叉的开销不划算,循环便保持串行。
在每个种子内部,加权均值的计算被写成矩阵-向量乘积;当种子这一维本身就已经把线程池填满时,实现会有意让这些计算保持串行,以免嵌套的 Rayon 分叉互相争抢。predict 在同一道门槛下并行它的最近中心扫描,判据是样本数乘簇数乘特征数。
对大多数数据集,你不用碰任何设置就能自动获得并行的好处;这些调优旋钮和阈值背后的思路,都在 7.3. 性能调优与并行 里。
2.9.7. 可复现性与持久化
拟合一个 MeanShift 是确定性的,不需要随机种子。默认播种时,它从每个点出发;bin seeding 时,它从每个格子固定的网格代表点出发。这两种方式都不从随机数生成器取值,所以在相同数据上做两次拟合,得到的中心和标签逐字节一致。
相比 KMeans,这是实打实的方便------KMeans 的质心初始化是随机的,要复现就得给种子。更全面的讨论见 7.1. 可复现性与随机种子。
这个模块里唯一引入随机性的地方是 estimate_bandwidth,而且只在你要它做子采样时才会引入。如果 n_samples 小于数据集,estimate_bandwidth 会打乱索引来挑选子集;当你需要估计本身可复现时,传一个固定的 random_state。请求全部行(n_samples 的默认值)会彻底消除随机性,因为根本没有可采样的余地。
拟合好的模型通过 save_to_path 序列化成一份紧凑的 postcard 二进制,用 load_from_path 还原。保存的文件带着中心、标签、超参数和训练元数据,因此重新加载的模型无需再次拟合就能给出完全一致的预测。
rust
use ndarray::Array2;
use rustyml::machine_learning::MeanShift;
fn main() {
let data = Array2::from_shape_vec(
(6, 2),
vec![0.0, 0.0, 0.1, 0.1, -0.1, 0.0, 10.0, 10.0, 10.1, 9.9, 9.9, 10.0],
)
.unwrap();
let mut ms = MeanShift::new(2.0).unwrap();
ms.fit(&data).unwrap();
let path = "mean_shift_model.bin";
ms.save_to_path(path).unwrap();
let restored = MeanShift::load_from_path(path).unwrap();
let before = ms.predict(&data).unwrap();
let after = restored.predict(&data).unwrap();
assert_eq!(before, after); // 往返一趟之后完全一致
std::fs::remove_file(path).unwrap();
println!("round-trip predictions match");
}
有一点要注意:如果你手上有核函数与合并规则变更之前落盘的模型文件,里面的中心来自旧的高斯核那一版。重新加载出的模型,对同一份数据的聚类结果会和重新拟合不一样------而且是悄无声息地不一样,文件本身并不会说明它是哪套算法产出的。请重新拟合、重新保存。
当拟合代价高、数据又稳定时,就把模型持久化下来,这样下游服务就能廉价地加载并预测。格式细节和版本兼容的注意事项,都在 7.2. 深入模型持久化 里。