【RustyML入门】2.9. MeanShift

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

非法参数会从 newwith_max_iterwith_toleranceError::InvalidParameter 的形式返回,正是 1.6. 错误处理 里讲过的那个错误类型。

2.9.3. 拟合、预测与读取结果

fit 接收一个二维数组,每行一个样本,跑完算法后返回 &mut Selfpredict 把新的点映射到学到的中心上,返回 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_bandwidthget_max_iterationsget_toleranceget_bin_seedingget_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. 深入模型持久化 里。

相关推荐
萧瑟其中~1 小时前
多线程锁详解:互斥锁·自旋锁·读写锁(CAS + futex 原理)
开发语言·c++
2401_894915531 小时前
GEO 源码部署如何实现精准地域分发?核心配置参数深度讲解
java·运维·服务器·后端·开源
IT_陈寒1 小时前
为什么我的Java Stream流操作会吃掉内存?
前端·人工智能·后端
初级代码游戏1 小时前
iOS开发 Swift 速记7:结构体和类
开发语言·ios·swift
golang学习记2 小时前
Go 项目使用docker compose的正确方式
开发语言·docker·golang
Java技术小馆2 小时前
LangChain 概述与生态
后端
用户250694921613 小时前
优雅的数据隔离:PostgreSQL 行级安全(RLS)
后端
Dr.kangder3 小时前
嵌入式总线设备解析——TTE总线应用与实践
开发语言·网络·算法·嵌入式·多任务·同步机制