用 Rust 从头实现机器学习算法·聚类①:k 均值、高斯混合模型与 EM

109 分钟阅读 rust-ml-from-scratch · 15
Rust机器学习

用 Rust 从头实现机器学习算法·聚类①:k 均值、高斯混合模型与 EM

聚类是无监督学习的另一副面孔:没有标签,只凭相似度让数据自己抱团,要求组内相似、组间相异。本篇是一对经典搭档——k 均值(硬分配的开山怪)与高斯混合模型(软分配的概率版),以及它们共同的引擎”迭代两步走”。

核心思想

k 均值的答案最朴素:簇内距离和(平方误差)最小。算法(Lloyd)反复执行两步:① 分配——每个点归到最近的质心;② 更新——每个质心移到簇内均值。两步都只会让目标下降,所以必然收敛——只是收敛到局部最优,因此初始化重要(本文用 k-means++ 风格的按距离平方概率选点)。

高斯混合模型(GMM)把问题概率化:数据由 KK 个二维高斯按权重 πk\pi_k 混合而成,p(x)=∑kπkN(x∣μk,Σk)p(x) = \sum_k \pi_k \mathcal{N}(x \mid \mu_k, \Sigma_k)。每个样本不再非此即彼,而是以责任度 γ(zik)\gamma(z_{ik}) 软属于各成分——重叠区里的点诚实地”脚踩两只船”。这正是清单的对照:k 均值输出硬标签,GMM 输出概率归属与不确定性量化。

清单还点破了 lineage:k 均值是 GMM-EM 的硬分配特例——把 GMM 的协方差固定为球形、责任度退化成 0/1,EM 就退化成 Lloyd 两步。

数学:EM 两步与单调性

EM(期望最大化)处理”有隐变量”的极大似然:

  • E 步:给定参数,算每个样本属于各成分的后验 γ(zik)=πkN(xi∣μk,Σk)∑jπjN(xi∣μj,Σj)\gamma(z_{ik}) = \frac{\pi_k \mathcal{N}(x_i \mid \mu_k, \Sigma_k)}{\sum_j \pi_j \mathcal{N}(x_i \mid \mu_j, \Sigma_j)};
  • M 步:把 γ\gamma 当软标签,加权重估参数:μk=∑iγikxiNk\mu_k = \frac{\sum_i \gamma_{ik} x_i}{N_k},Σk=∑iγik(xi−μk)(xi−μk)⊤Nk\Sigma_k = \frac{\sum_i \gamma_{ik}(x_i - \mu_k)(x_i - \mu_k)^{\top}}{N_k},πk=Nkn\pi_k = \frac{N_k}{n}。

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, &centroids, &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(&centroids, "#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) 是分配步的核心:每个点到 KK 个质心取最近。
  • Gmm<const K: usize> 用 const 泛型把成分数编进类型(语法角展开);协方差以 (a, b, c) 三元组存 2×22\times2 对称阵,行列式与逆都有闭式解——高斯密度不需要矩阵库。
  • E 步的 γ\gamma 计算在对数空间做:ln⁡πk+ln⁡N(xi)\ln \pi_k + \ln \mathcal{N}(x_i) 经 log-sum-exp 归一化,几百个指数相乘的场合不再下溢。
  • M 步每步给协方差加 10−410^{-4} 抖动,防退化簇把 Σ\Sigma 压奇异。
  • bic 实现”复杂度惩罚”:K=3K=3 的似然再好,参数罚款也让它出局(后文数据可见)。
// 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 &centroids {
            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 均值硬分配:球形边界斜切椭圆簇

GMM 软分配与 1σ 等高椭圆

EM 对数似然单调上升

怎么读这些数字和图

  • 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 与谱聚类,密度与图论两路奇兵。