用 Rust 从头实现机器学习算法·聚类①:k 均值、高斯混合模型与 EM
用 Rust 从头实现机器学习算法·聚类①:k 均值、高斯混合模型与 EM
聚类是无监督学习的另一副面孔:没有标签,只凭相似度让数据自己抱团,要求组内相似、组间相异。本篇是一对经典搭档——k 均值(硬分配的开山怪)与高斯混合模型(软分配的概率版),以及它们共同的引擎”迭代两步走”。
核心思想
k 均值的答案最朴素:簇内距离和(平方误差)最小。算法(Lloyd)反复执行两步:① 分配——每个点归到最近的质心;② 更新——每个质心移到簇内均值。两步都只会让目标下降,所以必然收敛——只是收敛到局部最优,因此初始化重要(本文用 k-means++ 风格的按距离平方概率选点)。
高斯混合模型(GMM)把问题概率化:数据由 个二维高斯按权重 混合而成,。每个样本不再非此即彼,而是以责任度 软属于各成分——重叠区里的点诚实地”脚踩两只船”。这正是清单的对照:k 均值输出硬标签,GMM 输出概率归属与不确定性量化。
清单还点破了 lineage:k 均值是 GMM-EM 的硬分配特例——把 GMM 的协方差固定为球形、责任度退化成 0/1,EM 就退化成 Lloyd 两步。
数学:EM 两步与单调性
EM(期望最大化)处理”有隐变量”的极大似然:
- E 步:给定参数,算每个样本属于各成分的后验 ;
- M 步:把 当软标签,加权重估参数:,,。
EM 的理论保证:每轮迭代似然不减(清单 4.2),收敛到局部最优。本文的输出曲线就是这个定理的实证。
Rust 实现
数据是两个斜向的椭圆簇(A 沿 +45° 拉长、B 沿 −45°,互相重叠)——球形假设(k 均值)在此吃亏,协方差假设(GMM)占优。
// src/main.rs(段一:k 均值、GMM/EM 与实验主程序)
// 聚类①:k 均值、高斯混合模型与 EM——数据自己抱团
// 单文件、仅标准库。GMM 用 const 泛型定成分数;绘图器与系列前篇一致。
// ===================== xorshift64 + Box-Muller(沿用 07 篇) =====================
struct XorShift(u64);
impl XorShift {
fn new(seed: u64) -> Self {
assert!(seed != 0, "种子不能为 0");
XorShift(seed)
}
fn next_u64(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn next_f64(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
}
fn next_gauss(&mut self) -> f64 {
let u1 = 1.0 - self.next_f64();
let u2 = self.next_f64();
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
}
}
// ===================== 2×2 特征分解(闭式,沿用 14 篇) =====================
fn eig2(a: f64, b: f64, c: f64) -> (f64, f64, (f64, f64), (f64, f64)) {
let t = a + c;
let d = a - c;
let disc = (d * d + 4.0 * b * b).sqrt();
let (mut l1, mut l2) = ((t + disc) / 2.0, (t - disc) / 2.0);
let (v1, v2);
if l1 < l2 {
std::mem::swap(&mut l1, &mut l2);
}
if b.abs() < 1e-12 {
v1 = if a >= c { (1.0, 0.0) } else { (0.0, 1.0) };
v2 = (-v1.1, v1.0);
} else {
let u1 = (b, l1 - a);
let n1 = (u1.0 * u1.0 + u1.1 * u1.1).sqrt();
v1 = (u1.0 / n1, u1.1 / n1);
v2 = (-v1.1, v1.0);
}
(l1, l2, v1, v2)
}
/// 2×2 协方差的行列式与逆。
fn det_inv2(a: f64, b: f64, c: f64) -> (f64, (f64, f64, f64)) {
let det = a * c - b * b;
assert!(det > 1e-12, "协方差矩阵应正定");
(det, (c / det, -b / det, a / det))
}
fn logsumexp(xs: &[f64]) -> f64 {
let m = xs.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
m + xs.iter().map(|x| (x - m).exp()).sum::<f64>().ln()
}
// ===================== 数据:两个斜向的椭圆簇(互相重叠) =====================
/// 簇 A 沿 +45° 拉长,簇 B 沿 −45° 拉长,中心 (0,0) 与 (2.5,2.5)。
/// 球形假设(k 均值)会在这里吃亏,协方差假设(GMM)占优。
fn make_data(rng: &mut XorShift, n: usize) -> Vec<(f64, f64)> {
let mut xs = Vec::with_capacity(n * 2);
for &(cx, cy, theta) in &[(0.0f64, 0.0, std::f64::consts::FRAC_PI_4), (2.5, 2.5, -std::f64::consts::FRAC_PI_4)] {
let (ct, st) = (theta.cos(), theta.sin());
for _ in 0..n {
let along = 2.0 * rng.next_gauss();
let across = 0.4 * rng.next_gauss();
xs.push((cx + along * ct - across * st, cy + along * st + across * ct));
}
}
xs
}
// ===================== k 均值(Lloyd 算法) =====================
fn dist2(a: (f64, f64), b: (f64, f64)) -> f64 {
(a.0 - b.0).powi(2) + (a.1 - b.1).powi(2)
}
/// Lloyd 两步迭代。返回 (质心, 分配, 惯性迭代曲线)。
fn kmeans(xs: &[(f64, f64)], k: usize, iters: usize, rng: &mut XorShift) -> (Vec<(f64, f64)>, Vec<usize>, Vec<(f64, f64)>) {
// k-means++ 风格初始化:先随机一点,再按距离平方比例选后续点
let mut centroids: Vec<(f64, f64)> = vec![xs[(rng.next_f64() * xs.len() as f64) as usize]];
while centroids.len() < k {
let d2: Vec<f64> = xs
.iter()
.map(|x| centroids.iter().map(|c| dist2(*x, *c)).fold(f64::INFINITY, f64::min))
.collect();
let total: f64 = d2.iter().sum();
let mut u = rng.next_f64() * total;
let mut pick = 0;
for (i, d) in d2.iter().enumerate() {
u -= d;
if u <= 0.0 {
pick = i;
break;
}
}
centroids.push(xs[pick]);
}
let mut assign = vec![0usize; xs.len()];
let mut curve = Vec::new();
for it in 1..=iters {
// 分配步
for (i, x) in xs.iter().enumerate() {
assign[i] = (0..k)
.min_by(|&a, &b| {
dist2(*x, centroids[a]).partial_cmp(&dist2(*x, centroids[b])).unwrap()
})
.unwrap();
}
// 更新步
let mut sum = vec![(0.0, 0.0); k];
let mut cnt = vec![0usize; k];
for (x, &a) in xs.iter().zip(assign.iter()) {
sum[a].0 += x.0;
sum[a].1 += x.1;
cnt[a] += 1;
}
for c in 0..k {
if cnt[c] > 0 {
centroids[c] = (sum[c].0 / cnt[c] as f64, sum[c].1 / cnt[c] as f64);
}
}
let inertia: f64 = xs.iter().zip(assign.iter()).map(|(x, &a)| dist2(*x, centroids[a])).sum();
curve.push((it as f64, inertia));
if it >= 3 && (curve[it - 2].1 - inertia).abs() < 1e-9 * inertia.abs().max(1.0) {
break;
}
}
(centroids, assign, curve)
}
/// 轮廓系数:s(i) = (b − a) / max(a, b),越接近 1 聚得越紧。
fn silhouette(xs: &[(f64, f64)], assign: &[usize], k: usize) -> f64 {
let n = xs.len();
let mut scores = 0.0;
for i in 0..n {
let mut a = 0.0;
let mut b = f64::INFINITY;
for c in 0..k {
let mut sum = 0.0;
let mut cnt = 0usize;
for j in 0..n {
if assign[j] == c {
sum += dist2(xs[i], xs[j]).sqrt();
cnt += 1;
}
}
if c == assign[i] {
a = if cnt > 1 { sum / (cnt - 1) as f64 } else { 0.0 };
} else if cnt > 0 {
b = b.min(sum / cnt as f64);
}
}
let s = if a.max(b) > 0.0 { (b - a) / a.max(b) } else { 0.0 };
scores += s;
}
scores / n as f64
}
// ===================== 高斯混合模型与 EM(const 泛型定成分数) =====================
/// K 个二维高斯的混合。μ: [[f64; 2]; K],Σ 以 (a, b, c) 三元组存 [[a,b],[b,c]]。
struct Gmm<const K: usize> {
weights: [f64; K],
mu: [[f64; 2]; K],
sigma: [(f64, f64, f64); K],
}
impl<const K: usize> Gmm<K> {
/// 用 k 均值的质心与簇协方差做 EM 初始化(标准实践)。
fn init_from_kmeans(xs: &[(f64, f64)], centroids: &[(f64, f64)], assign: &[usize]) -> Self {
let n = xs.len() as f64;
let mut mu = [[0.0; 2]; K];
let mut sigma = [(0.0, 0.0, 0.0); K];
let mut weights = [0.0; K];
for k in 0..K {
let pts: Vec<&(f64, f64)> = xs.iter().zip(assign).filter(|(_, &a)| a == k).map(|(x, _)| x).collect();
let m = pts.len().max(1) as f64;
weights[k] = m / n;
mu[k] = [
pts.iter().map(|p| p.0).sum::<f64>() / m,
pts.iter().map(|p| p.1).sum::<f64>() / m,
];
let (mut a, mut b, mut c) = (0.0, 0.0, 0.0);
for p in &pts {
let dx = p.0 - mu[k][0];
let dy = p.1 - mu[k][1];
a += dx * dx;
b += dx * dy;
c += dy * dy;
}
sigma[k] = (a / m + 1e-4, b / m, c / m + 1e-4); // 抖动防奇异
}
let _ = centroids;
Gmm { weights, mu, sigma }
}
/// 单点的高斯对数密度(二维)。
fn log_gauss(&self, x: (f64, f64), k: usize) -> f64 {
let (a, b, c) = self.sigma[k];
let (det, (ia, ib, ic)) = det_inv2(a, b, c);
let dx = x.0 - self.mu[k][0];
let dy = x.1 - self.mu[k][1];
let d2 = ia * dx * dx + 2.0 * ib * dx * dy + ic * dy * dy;
-0.5 * (2.0 * std::f64::consts::TAU.ln() + det.ln() + d2)
}
/// 数据的对数似然 Σ ln Σ_k π_k N(x_i | μ_k, Σ_k)。
fn log_likelihood(&self, xs: &[(f64, f64)]) -> f64 {
xs.iter()
.map(|x| {
let terms: Vec<f64> = (0..K)
.map(|k| (self.weights[k]).ln() + self.log_gauss(*x, k))
.collect();
logsumexp(&terms)
})
.sum()
}
/// EM 迭代:返回 (责任度 γ, 对数似然迭代曲线)。
fn fit(&mut self, xs: &[(f64, f64)], iters: usize) -> (Vec<[f64; K]>, Vec<(f64, f64)>) {
let mut curve = Vec::new();
let mut gamma = vec![[0.0; K]; xs.len()];
for it in 1..=iters {
// E 步:γ_ik ∝ π_k · N(x_i | μ_k, Σ_k),用 log-sum-exp 防下溢
for (i, x) in xs.iter().enumerate() {
let terms: Vec<f64> = (0..K)
.map(|k| (self.weights[k]).ln() + self.log_gauss(*x, k))
.collect();
let norm = logsumexp(&terms);
for k in 0..K {
gamma[i][k] = ((terms[k]) - norm).exp();
}
}
// M 步:按责任度加权重估参数
let n = xs.len() as f64;
for k in 0..K {
let nk: f64 = gamma.iter().map(|g| g[k]).sum();
self.weights[k] = nk / n;
self.mu[k] = [0.0; 2];
for (x, g) in xs.iter().zip(gamma.iter()) {
self.mu[k][0] += g[k] * x.0;
self.mu[k][1] += g[k] * x.1;
}
self.mu[k][0] /= nk;
self.mu[k][1] /= nk;
let (mut a, mut b, mut c) = (0.0, 0.0, 0.0);
for (x, g) in xs.iter().zip(gamma.iter()) {
let dx = x.0 - self.mu[k][0];
let dy = x.1 - self.mu[k][1];
a += g[k] * dx * dx;
b += g[k] * dx * dy;
c += g[k] * dy * dy;
}
self.sigma[k] = (a / nk + 1e-4, b / nk, c / nk + 1e-4);
}
curve.push((it as f64, self.log_likelihood(xs)));
if it >= 3 {
let gain = curve[it - 1].1 - curve[it - 2].1;
if gain.abs() < 1e-6 * curve[it - 1].1.abs().max(1.0) {
break;
}
}
}
(gamma, curve)
}
/// BIC = −2·LL + 参数数·ln n,用于选成分数(惩罚复杂度)。
fn bic(&self, xs: &[(f64, f64)]) -> f64 {
let params = K * 6 - 1; // 每成分:均值 2 + 协方差 3 + 权重 1,权重和为 1 减 1
-2.0 * self.log_likelihood(xs) + params as f64 * (xs.len() as f64).ln()
}
}
fn main() {
let out_dir = "../../../frontend/public/images/series/rust-ml-15-kmeans-gmm";
std::fs::create_dir_all(out_dir).expect("创建输出目录失败");
// ---- 1. 数据 ----
let mut rng = XorShift::new(42);
let xs = make_data(&mut rng, 120);
println!("数据:{} 个点,两个斜向椭圆簇(A 沿 +45°,B 沿 −45°,互相重叠)", xs.len());
// ---- 2. k 均值 ----
let (centroids, assign, km_curve) = kmeans(&xs, 2, 100, &mut rng);
let inertia = km_curve.last().unwrap().1;
println!(
"\n[k 均值] 质心 ≈ ({:.2}, {:.2}) 与 ({:.2}, {:.2}),惯性 {:.1},轮廓系数 {:.3},收敛于 {} 轮",
centroids[0].0, centroids[0].1,
centroids[1].0, centroids[1].1,
inertia,
silhouette(&xs, &assign, 2),
km_curve.len()
);
// ---- 3. GMM + EM(k=2 与 BIC 选成分数) ----
let mut gmm = Gmm::<2>::init_from_kmeans(&xs, ¢roids, &assign);
let (gamma, ll_curve) = gmm.fit(&xs, 100);
println!(
"[GMM k=2] 对数似然 {:.2},收敛于 {} 轮;权重 [{:.2}, {:.2}]",
ll_curve.last().unwrap().1,
ll_curve.len(),
gmm.weights[0],
gmm.weights[1]
);
println!("\n[EM 对数似然(单调性证据)] 前 8 轮:");
for (it, ll) in ll_curve.iter().take(8) {
println!(" 轮 {:>2}:{:.4}", it, ll);
}
// BIC 选成分数
println!("\n[BIC 选成分数](越小越好)");
{
let (c1, a1) = { let (c, a, _) = kmeans(&xs, 1, 50, &mut rng); (c, a) };
let mut g1 = Gmm::<1>::init_from_kmeans(&xs, &c1, &a1);
g1.fit(&xs, 50);
println!(" K=1:BIC = {:.1}", g1.bic(&xs));
println!(" K=2:BIC = {:.1}", gmm.bic(&xs));
let (c3, a3) = { let (c, a, _) = kmeans(&xs, 3, 50, &mut rng); (c, a) };
let mut g3 = Gmm::<3>::init_from_kmeans(&xs, &c3, &a3);
g3.fit(&xs, 50);
println!(" K=3:BIC = {:.1}", g3.bic(&xs));
}
// ---- 4. 图一:k 均值硬分配 ----
let cluster_a: Vec<(f64, f64)> = xs.iter().zip(assign.iter()).filter(|(_, &a)| a == 0).map(|(x, _)| *x).collect();
let cluster_b: Vec<(f64, f64)> = xs.iter().zip(assign.iter()).filter(|(_, &a)| a == 1).map(|(x, _)| *x).collect();
let mut c = Canvas::new(560.0, 460.0, -5.0, 7.0, -5.0, 7.0);
c.axes("x1", "x2");
c.dots(&cluster_a, PALETTE[1], 3.0);
c.dots(&cluster_b, PALETTE[3], 3.0);
c.dots(¢roids, "#16161d", 6.0);
c.legend(&[("簇 0", PALETTE[1]), ("簇 1", PALETTE[3]), ("质心", "#16161d")]);
let p1 = format!("{out_dir}/kmeans.svg");
c.save(&p1);
// ---- 5. 图二:GMM 软分配 + 高斯等高椭圆 ----
let soft_a: Vec<(f64, f64)> = xs.iter().zip(gamma.iter()).filter(|(_, g)| g[0] >= g[1]).map(|(x, _)| *x).collect();
let soft_b: Vec<(f64, f64)> = xs.iter().zip(gamma.iter()).filter(|(_, g)| g[0] < g[1]).map(|(x, _)| *x).collect();
let mut c = Canvas::new(560.0, 460.0, -5.0, 7.0, -5.0, 7.0);
c.axes("x1", "x2");
c.dots(&soft_a, PALETTE[1], 3.0);
c.dots(&soft_b, PALETTE[3], 3.0);
// 每个成分画 1σ 椭圆:点 = μ + V·diag(√λ)·(cos t, sin t)
for k in 0..2 {
let (a, b, cc) = gmm.sigma[k];
let (l1, l2, v1, _) = eig2(a, b, cc);
let (s1, s2) = (l1.sqrt(), l2.sqrt());
let ellipse: Vec<(f64, f64)> = (0..=80)
.map(|i| {
let t = i as f64 / 80.0 * std::f64::consts::TAU;
let (ex, ey) = (s1 * t.cos(), s2 * t.sin());
(gmm.mu[k][0] + ex * v1.0 - ey * v1.1, gmm.mu[k][1] + ex * v1.1 + ey * v1.0)
})
.collect();
c.polyline(&ellipse, PALETTE[0], 2.0);
}
c.dots(&gmm.mu.iter().map(|m| (m[0], m[1])).collect::<Vec<_>>(), "#16161d", 5.0);
c.legend(&[("成分 0 为主", PALETTE[1]), ("成分 1 为主", PALETTE[3]), ("1σ 椭圆", PALETTE[0])]);
let p2 = format!("{out_dir}/gmm.svg");
c.save(&p2);
// ---- 6. 图三:EM 对数似然单调上升 ----
let mut c = Canvas::new(560.0, 320.0, 1.0, ll_curve.len() as f64, 0.0, 0.0);
c.axes("EM 迭代轮数", "对数似然");
c.polyline(&ll_curve, PALETTE[0], 2.0);
let p3 = format!("{out_dir}/em-likelihood.svg");
c.save(&p3);
println!("图已写入:{p1}、{p2}、{p3}");
}
// ===================== 迷你 SVG 绘图器(与系列前篇一致) =====================
实现要点:
kmeans里.min_by(partial_cmp)是分配步的核心:每个点到 个质心取最近。Gmm<const K: usize>用 const 泛型把成分数编进类型(语法角展开);协方差以(a, b, c)三元组存 对称阵,行列式与逆都有闭式解——高斯密度不需要矩阵库。- E 步的 计算在对数空间做: 经 log-sum-exp 归一化,几百个指数相乘的场合不再下溢。
- M 步每步给协方差加 抖动,防退化簇把 压奇异。
bic实现”复杂度惩罚”: 的似然再好,参数罚款也让它出局(后文数据可见)。
// src/main.rs 段二:迷你 SVG 绘图器(与系列前篇同一份实现)
use std::fmt::Write as _;
/// 系列默认配色:朱橙 / 青绿 / 蓝 / 琥珀(取自站点设计令牌)。
const PALETTE: [&str; 4] = ["#d6491f", "#0f766e", "#2563eb", "#d97706"];
/// 一张图:坐标映射 + 已累积的 SVG 元素。
struct Canvas {
w: f64,
h: f64,
xmin: f64,
xmax: f64,
ymin: f64,
ymax: f64,
pad_l: f64,
pad_r: f64,
pad_t: f64,
pad_b: f64,
body: String,
}
impl Canvas {
/// 按数据范围自动留 5% 边距建图;xmin == xmax 时手动摊开避免除零。
fn new(w: f64, h: f64, xmin: f64, xmax: f64, ymin: f64, ymax: f64) -> Canvas {
let (xmin, xmax) = spread(xmin, xmax);
let (ymin, ymax) = spread(ymin, ymax);
Canvas {
w,
h,
xmin,
xmax,
ymin,
ymax,
pad_l: 52.0,
pad_r: 16.0,
pad_t: 16.0,
pad_b: 40.0,
body: String::new(),
}
}
fn px(&self, x: f64) -> f64 {
self.pad_l + (x - self.xmin) / (self.xmax - self.xmin) * (self.w - self.pad_l - self.pad_r)
}
fn py(&self, y: f64) -> f64 {
self.h - self.pad_b - (y - self.ymin) / (self.ymax - self.ymin) * (self.h - self.pad_t - self.pad_b)
}
/// 折线,点需按 x 升序传入。
fn polyline(&mut self, pts: &[(f64, f64)], color: &str, width: f64) {
let mut d = String::new();
for (i, (x, y)) in pts.iter().enumerate() {
let _ = write!(d, "{}{:.1},{:.1}", if i == 0 { "M" } else { "L" }, self.px(*x), self.py(*y));
}
let _ = write!(
self.body,
r##"<path d="{d}" fill="none" stroke="{color}" stroke-width="{width}" stroke-linejoin="round"/>"##
);
}
/// 散点。
fn dots(&mut self, pts: &[(f64, f64)], color: &str, r: f64) {
for (x, y) in pts {
let _ = write!(
self.body,
r##"<circle cx="{:.1}" cy="{:.1}" r="{r}" fill="{color}" fill-opacity="0.65"/>"##,
self.px(*x),
self.py(*y)
);
}
}
/// 坐标轴 + 虚线网格 + 刻度标签。
fn axes(&mut self, xlabel: &str, ylabel: &str) {
let x0 = self.px(self.xmin);
let x1 = self.px(self.xmax);
let y0 = self.py(self.ymin);
let y1 = self.py(self.ymax);
for t in nice_ticks(self.xmin, self.xmax, 6) {
let x = self.px(t);
let _ = write!(
self.body,
r##"<line x1="{x:.1}" y1="{y0:.1}" x2="{x:.1}" y2="{y1:.1}" stroke="#16161d" stroke-opacity="0.08"/>"##
);
let _ = write!(
self.body,
r##"<text x="{x:.1}" y="{:.1}" text-anchor="middle" fill="#16161d" fill-opacity="0.45">{t:.6}</text>"##,
y0 + 16.0,
t = trim(t)
);
}
for t in nice_ticks(self.ymin, self.ymax, 6) {
let y = self.py(t);
let _ = write!(
self.body,
r##"<line x1="{x0:.1}" y1="{y:.1}" x2="{x1:.1}" y2="{y:.1}" stroke="#16161d" stroke-opacity="0.08"/>"##
);
let _ = write!(
self.body,
r##"<text x="{:.1}" y="{:.1}" text-anchor="end" fill="#16161d" fill-opacity="0.45">{t:.6}</text>"##,
x0 - 6.0,
y + 4.0,
t = trim(t)
);
}
let _ = write!(
self.body,
r##"<path d="M{x0:.1},{y0:.1}H{x1:.1}M{x0:.1},{y0:.1}V{y1:.1}" fill="none" stroke="#16161d" stroke-opacity="0.6"/>"##
);
let _ = write!(
self.body,
r##"<text x="{:.1}" y="{:.1}" text-anchor="middle" fill="#16161d" fill-opacity="0.7">{xlabel}</text>"##,
(x0 + x1) / 2.0,
self.h - 8.0
);
let _ = write!(
self.body,
r##"<text x="14" y="{:.1}" text-anchor="middle" transform="rotate(-90 14 {:.1})" fill="#16161d" fill-opacity="0.7">{ylabel}</text>"##,
(y0 + y1) / 2.0,
(y0 + y1) / 2.0
);
}
/// 置信带:upper/lower 两条曲线围成的区域,半透明填充。
#[allow(dead_code)] // 部分篇章不使用,保持各篇绘图器逐字一致
fn band(&mut self, upper: &[(f64, f64)], lower: &[(f64, f64)], color: &str) {
let mut d = String::new();
for (i, (x, y)) in upper.iter().enumerate() {
let _ = write!(d, "{}{:.1},{:.1}", if i == 0 { "M" } else { "L" }, self.px(*x), self.py(*y));
}
for (x, y) in lower.iter().rev() {
let _ = write!(d, "L{:.1},{:.1}", self.px(*x), self.py(*y));
}
let _ = write!(
self.body,
r##"<path d="{d}Z" fill="{color}" fill-opacity="0.15" stroke="none"/>"##
);
}
/// 图例:右上角依次画「色线 + 文字」。
fn legend(&mut self, entries: &[(&str, &str)]) { let sample_w = 22.0;
let line_h = 18.0;
let x_text = self.w - self.pad_r - 78.0 + sample_w + 6.0;
for (i, (label, color)) in entries.iter().enumerate() {
let y = self.pad_t + 14.0 + i as f64 * line_h;
let x0 = self.w - self.pad_r - 78.0;
let _ = write!(
self.body,
r##"<line x1="{x0:.1}" y1="{y:.1}" x2="{:.1}" y2="{y:.1}" stroke="{color}" stroke-width="2.5"/>"##,
x0 + sample_w
);
let _ = write!(
self.body,
r##"<text x="{x_text:.1}" y="{:.1}" fill="#16161d" fill-opacity="0.7">{label}</text>"##,
y + 4.0
);
}
}
fn save(&self, path: &str) {
let svg = format!(
r##"<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 {w} {h}" font-family="'JetBrains Mono',monospace" font-size="11"><rect width="{w}" height="{h}" fill="#fbfaf7"/>{body}</svg>"##,
w = self.w,
h = self.h,
body = self.body
);
std::fs::write(path, svg).expect("写 SVG 失败");
}
}
/// 把 fmt 出来的 f64 末尾零去掉,让刻度标签短一点。
fn trim(v: f64) -> String {
let s = format!("{v:.6}");
s.trim_end_matches('0').trim_end_matches('.').to_string()
}
/// 数据范围退化成一点时向两侧摊开。
fn spread(min: f64, max: f64) -> (f64, f64) {
if min == max {
(min - 1.0, max + 1.0)
} else {
let pad = (max - min) * 0.05;
(min - pad, max + pad)
}
}
/// “好看”的刻度:步长取 1/2/2.5/5 × 10^k。
fn nice_ticks(min: f64, max: f64, n: usize) -> Vec<f64> {
let raw = (max - min) / n.max(1) as f64;
if raw <= 0.0 || !raw.is_finite() {
return vec![min];
}
let exp = raw.abs().log10().floor() as i32;
let base = 10f64.powi(exp);
let frac = raw / base;
let step = if frac <= 1.0 {
base
} else if frac <= 2.0 {
2.0 * base
} else if frac <= 2.5 {
2.5 * base
} else if frac <= 5.0 {
5.0 * base
} else {
10.0 * base
};
let mut ticks = Vec::new();
let mut t = (min / step).ceil() * step;
while t <= max + 1e-9 {
ticks.push(if t.abs() < step * 1e-9 { 0.0 } else { t });
t += step;
}
ticks
}
// src/main.rs 末尾:单元测试(cargo test 运行)
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn kmeans_recovers_centers() {
let mut rng = XorShift::new(7);
let xs = make_data(&mut rng, 80);
let (centroids, _, _) = kmeans(&xs, 2, 100, &mut XorShift::new(3));
// 质心应接近 (0,0) 与 (2.5,2.5)(顺序不定)
// 球形假设在斜椭圆簇上会向重叠区偏移,容差放宽到 1.2
let mut near = [false; 2];
for c in ¢roids {
if (c.0 - 0.0).abs() < 1.2 && (c.1 - 0.0).abs() < 1.2 {
near[0] = true;
}
if (c.0 - 2.5).abs() < 1.2 && (c.1 - 2.5).abs() < 1.2 {
near[1] = true;
}
}
assert!(near[0] && near[1], "质心应落在真实中心附近: {centroids:?}");
}
#[test]
fn em_log_likelihood_is_monotone() {
let mut rng = XorShift::new(11);
let xs = make_data(&mut rng, 60);
let mut rng2 = XorShift::new(3);
let (c, a, _) = kmeans(&xs, 2, 50, &mut rng2);
let mut gmm = Gmm::<2>::init_from_kmeans(&xs, &c, &a);
let (_, curve) = gmm.fit(&xs, 60);
for w in curve.windows(2) {
assert!(w[1].1 >= w[0].1 - 1e-9, "对数似然应单调不减: {} → {}", w[0].1, w[1].1);
}
}
#[test]
fn responsibilities_are_normalized() {
let mut rng = XorShift::new(13);
let xs = make_data(&mut rng, 50);
let mut rng2 = XorShift::new(3);
let (c, a, _) = kmeans(&xs, 2, 50, &mut rng2);
let mut gmm = Gmm::<2>::init_from_kmeans(&xs, &c, &a);
let (gamma, _) = gmm.fit(&xs, 30);
for g in &gamma {
assert!((g[0] + g[1] - 1.0).abs() < 1e-9);
}
}
#[test]
fn logsumexp_is_stable() {
// 极端值:直接 exp 会下溢,log-sum-exp 不会
let xs = [-1000.0, -1001.0, -999.0];
let v = logsumexp(&xs);
assert!(v.is_finite() && (v - -998.593).abs() < 0.01);
}
}
Rust 语法角:const 泛型——把常数编进类型
struct Gmm<const K: usize> 声明了一个”成分数为编译期常数”的类型:Gmm<2> 与 Gmm<3> 是不同的类型,它们的 weights: [f64; K] 等数组长度由类型系统担保。对照 Python:NumPy 里数组形状是运行值,传错成分数要到运行时才发现;Rust 里 Gmm<2>::init_from_kmeans 想塞三个质心是编译错误。main 里对 K=1/2/3 各跑一遍 BIC,正是 const 泛型的顺手用法。见《Rust 程序设计语言》ch10-01(泛型)。
运行结果
cargo test(4 个用例:质心恢复、对数似然单调不减、责任度归一、log-sum-exp 稳定性)全部通过后,cargo run:
数据:240 个点,两个斜向椭圆簇(A 沿 +45°,B 沿 −45°,互相重叠)
[k 均值] 质心 ≈ (-0.50, -0.53) 与 (2.21, 2.48),惯性 702.7,轮廓系数 0.525,收敛于 6 轮
[GMM k=2] 对数似然 -761.50,收敛于 12 轮;权重 [0.50, 0.50]
[EM 对数似然(单调性证据)] 前 8 轮:
轮 1:-771.3328
轮 2:-764.6061
轮 3:-762.4664
轮 4:-761.8486
轮 5:-761.6357
轮 6:-761.5550
轮 7:-761.5235
轮 8:-761.5110
[BIC 选成分数](越小越好)
K=1:BIC = 1930.7
K=2:BIC = 1583.3
K=3:BIC = 1609.7
图已写入:../../../frontend/public/images/series/rust-ml-15-kmeans-gmm/kmeans.svg、../../../frontend/public/images/series/rust-ml-15-kmeans-gmm/gmm.svg、../../../frontend/public/images/series/rust-ml-15-kmeans-gmm/em-likelihood.svg
三张图:
怎么读这些数字和图
- k 均值质心的偏移 (-0.50, -0.53) 就是球形假设的代价:真实中心在 (0,0),但球形簇没法贴合斜椭圆,质心被重叠区的”中间质量”拖走。GMM 的均值则因协方差会”斜着长”而贴近真值——同一批数据,假设不同,答案质量立判。
- 对数似然曲线 −771.3 → −761.5:每轮只增不减,增益逐轮收窄——EM 的收敛定理在数字里显形。这也是清单说的”每次迭代保证似然不减”的直接证据。
- BIC 三选一:K=2 胜(1583.3 < 1930.7 < 1609.7):K=1 似然太差,K=3 用 6 个额外参数换的似然提升不够付罚款——没有标签,模型选择照样能做,这就是信息准则的价值。
- GMM 图上的 1σ 椭圆是”成分画像”:两个椭圆各自贴着斜向数据的长轴,重叠区里的点按较大的责任度着色——软分配的”脚踩两只船”在图上可见。k 均值的图里边界则是一条生硬的斜线,切掉了不少同簇的点。
优缺点与适用场景
(抄清单 4.1–4.3 原文)
- k 均值:简单快速、可解释、适合大样本;但须预设 k、只发现球形簇、对异常值与初始值敏感。选 k 用肘部法则、轮廓系数、Calinski-Harabasz 指数(本文实现了轮廓系数:0.525,中等偏好的聚类)。
- EM:处理隐变量极大似然的通用引擎,保证似然不减,收敛到局部最优;也用于缺失数据填补。
- GMM:软聚类(概率归属)、概率解释完整、可作密度估计器(低概率样本即异常);但需预设成分数、易收敛到局部最优、高维协方差参数多。适用场景:语音处理、图像分割、异常检测。
小结
k 均值与 GMM 是”迭代两步走”思维的双生子:分配-更新与 E-M 共享同一个骨架,差别只在”软”与”硬”。配 BIC 完成模型选择,无监督的流程闭环了。但它们的簇都是”凸的蛋形”——月牙形、环形的簇怎么办?下一篇聚类②:DBSCAN 与谱聚类,密度与图论两路奇兵。