用 Rust 从头实现机器学习算法·回归⑤:贝叶斯回归——让模型说出"我有多确定"

94 分钟阅读 rust-ml-from-scratch · 7
Rust机器学习

用 Rust 从头实现机器学习算法·回归⑤:贝叶斯回归——让模型说出”我有多确定”

前四篇的回归模型训练完只吐出一个答案:一条线、一组系数。它有多可靠?数据稀疏处和外推处的预测可信吗?这些问题在点估计框架里没有容身之处。本篇进入贝叶斯世界:参数不再是固定的未知数,而是随机变量——模型输出从”一个答案”变成”一个答案 + 一个误差棒”。

核心思想

清单里贝叶斯回归的定义一句话:先验分布 + 似然 → 后验分布,用整个后验做推断和预测,而不是只给一个点估计。它的操作序列是贝叶斯公式的直白翻译:

  1. 训练前,对参数 ww 有一个先验信念:w∼N(0,τ2I)w \sim \mathcal{N}(0, \tau^2 I)(“系数大概率不大”——听出来了吗,这就是 06 篇岭回归罚项的贝叶斯面孔);
  2. 看到数据后,按贝叶斯定理把先验更新成后验 w∣y∼N(μ,Σ)w \mid y \sim \mathcal{N}(\mu, \Sigma);
  3. 预测新样本时,对后验里所有可能的 ww 做加权平均——预测本身也是分布,天然携带不确定度。

后验还有个好性质:它是序贯的。每来一个新数据点,把当前后验当作先验再更新一次即可——“在线学习”不需要重写训练循环。

数学:共轭高斯的三行推导结果

模型设定:y=w⊤ϕ(x)+εy = w^{\top}\phi(x) + \varepsilon,ε∼N(0,σ2)\varepsilon \sim \mathcal{N}(0, \sigma^2)(噪声方差 σ2\sigma^2 当作已知,实践中从残差估计);先验 w∼N(0,τ2I)w \sim \mathcal{N}(0, \tau^2 I)。高斯先验 + 高斯似然是共轭对,后验仍是高斯,直接把精度矩阵(协方差的逆)写出来:

Σ−1=1σ2Φ⊤Φ+1τ2I,μ=Σ Φ⊤yσ2\Sigma^{-1} = \frac{1}{\sigma^2}\Phi^{\top}\Phi + \frac{1}{\tau^2}I, \qquad \mu = \Sigma\, \frac{\Phi^{\top} y}{\sigma^2}

对照 06 篇:后验均值 μ\mu 恰好最小化 ∥Φw−y∥2+σ2τ2∥w∥2\|\Phi w - y\|^2 + \frac{\sigma^2}{\tau^2}\|w\|^2——MAP 估计就是 λ=σ2/τ2\lambda = \sigma^2/\tau^2 的岭回归。06 篇的 λ 扫描,在贝叶斯语言里就是”先验强度扫描”。

对新特征 ϕ∗\phi_* 的预测分布(对后验积分后仍是高斯):

y∗∣y∼N(μ⊤ϕ∗, σ2⏟噪声+ϕ∗⊤Σ ϕ∗⏟参数不确定度)y_* \mid y \sim \mathcal{N}\big(\mu^{\top}\phi_*,\ \underbrace{\sigma^2}_{\text{噪声}} + \underbrace{\phi_*^{\top}\Sigma\,\phi_*}_{\text{参数不确定度}}\big)

方差的分解极具解释力:哪怕参数完全确定(Σ→0\Sigma \to 0),预测仍至少有 σ2\sigma^2 的噪声地板;而数据稀疏处 Σ\Sigma 大,误差棒自动张开——模型知道自己在外推。

Rust 实现

系列第一次出现”模型即结构体”:BayesReg 持有两个超参数,拟合与预测都是它的方法。

// src/main.rs(段一:Matrix、Box-Muller 采样、BayesReg 与主程序)
// 回归⑤:贝叶斯回归——把参数当作随机变量
// 单文件、仅标准库。绘图器与 01~06 篇内联的是同一份实现(本篇起新增 band 置信带)。

// ===================== Matrix(含高斯-约当求逆,沿用 06 篇) =====================

#[derive(Debug, Clone)]
struct Matrix {
    data: Vec<f64>,
    rows: usize,
    cols: usize,
}

impl Matrix {
    fn from_vec(data: Vec<f64>, rows: usize, cols: usize) -> Matrix {
        assert_eq!(data.len(), rows * cols, "数据长度与矩阵形状不一致");
        Matrix { data, rows, cols }
    }

    fn zeros(rows: usize, cols: usize) -> Matrix {
        Matrix { data: vec![0.0; rows * cols], rows, cols }
    }

    fn shape(&self) -> (usize, usize) {
        (self.rows, self.cols)
    }

    fn get(&self, i: usize, j: usize) -> f64 {
        self.data[i * self.cols + j]
    }

    fn matmul(&self, other: &Matrix) -> Matrix {
        assert_eq!(self.cols, other.rows, "内维不一致:{:?} 无法乘 {:?}", self.shape(), other.shape());
        let mut out = Matrix::zeros(self.rows, other.cols);
        for i in 0..self.rows {
            for j in 0..other.cols {
                let mut sum = 0.0;
                for k in 0..self.cols {
                    sum += self.get(i, k) * other.get(k, j);
                }
                out.data[i * out.cols + j] = sum;
            }
        }
        out
    }

    fn transpose(&self) -> Matrix {
        let mut out = Matrix::zeros(self.cols, self.rows);
        for i in 0..self.rows {
            for j in 0..self.cols {
                out.data[j * self.rows + i] = self.get(i, j);
            }
        }
        out
    }

    /// 高斯-约当消元求逆(带部分主元)。奇异矩阵返回 None。
    fn inverse(&self) -> Option<Matrix> {
        assert_eq!(self.rows, self.cols, "只有方阵才能求逆");
        let n = self.rows;
        let mut aug = Matrix::zeros(n, 2 * n);
        for i in 0..n {
            for j in 0..n {
                aug.data[i * 2 * n + j] = self.get(i, j);
            }
            aug.data[i * 2 * n + n + i] = 1.0;
        }
        for col in 0..n {
            let mut pivot = col;
            for r in col + 1..n {
                if aug.get(r, col).abs() > aug.get(pivot, col).abs() {
                    pivot = r;
                }
            }
            if aug.get(pivot, col).abs() < 1e-12 {
                return None;
            }
            if pivot != col {
                for j in 0..2 * n {
                    aug.data.swap(col * 2 * n + j, pivot * 2 * n + j);
                }
            }
            let p = aug.get(col, col);
            for j in 0..2 * n {
                aug.data[col * 2 * n + j] /= p;
            }
            for r in 0..n {
                if r == col {
                    continue;
                }
                let f = aug.get(r, col);
                if f == 0.0 {
                    continue;
                }
                for j in 0..2 * n {
                    aug.data[r * 2 * n + j] -= f * aug.get(col, j);
                }
            }
        }
        let mut inv = Matrix::zeros(n, n);
        for i in 0..n {
            for j in 0..n {
                inv.data[i * n + j] = aug.get(i, n + j);
            }
        }
        Some(inv)
    }
}

// ===================== xorshift64 + Box-Muller 正态采样 =====================

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
    }

    /// Box-Muller:两个独立 Uniform(0,1) → 一对独立标准正态。
    fn next_gauss(&mut self) -> f64 {
        let u1 = 1.0 - self.next_f64(); // 防 ln(0)
        let u2 = self.next_f64();
        (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
    }
}

// ===================== 贝叶斯线性回归(共轭高斯模型) =====================

/// 模型:y = wᵀφ(x) + ε,ε ~ N(0, σ²);先验 w ~ N(0, τ²I)。
/// 后验 w|y ~ N(μ, Σ),Σ = (ΦᵀΦ/σ² + I/τ²)⁻¹,μ = ΣΦᵀy/σ²。
struct BayesReg {
    sigma2: f64, // 噪声方差(已知假设)
    tau2: f64,   // 先验方差
}

impl BayesReg {
    fn new(sigma2: f64, tau2: f64) -> Self {
        assert!(sigma2 > 0.0 && tau2 > 0.0);
        BayesReg { sigma2, tau2 }
    }

    /// 对应岭回归的 λ:λ = σ²/τ²——先验越强(τ 越小)等价于罚得越重。
    fn ridge_lambda(&self) -> f64 {
        self.sigma2 / self.tau2
    }

    /// 拟合:返回 (后验均值 μ, 后验协方差 Σ)。
    fn fit(&self, phi: &Matrix, ys: &[f64]) -> (Matrix, Matrix) {
        let (n, cols) = phi.shape();
        let xt = phi.transpose();
        let mut a = xt.matmul(phi);
        for i in 0..a.data.len() {
            a.data[i] /= self.sigma2;
        }
        for j in 0..cols {
            a.data[j * cols + j] += 1.0 / self.tau2;
        }
        let sigma = a.inverse().expect("后验精度矩阵必正定可逆");
        let b = xt.matmul(&Matrix::from_vec(ys.to_vec(), n, 1));
        let mu = sigma.matmul(&b); // Σ·(Φᵀy/σ²)
        let mut mu = mu;
        for v in mu.data.iter_mut() {
            *v /= self.sigma2;
        }
        (mu, sigma)
    }

    /// 后验预测:均值 μᵀφ 与总方差 σ² + φᵀΣφ。
    fn predict(&self, phi_x: &Matrix, mu: &Matrix, sigma: &Matrix) -> (f64, f64) {
        let mean = phi_x.transpose().matmul(mu).data[0];
        let var = self.sigma2 + phi_x.transpose().matmul(sigma).matmul(phi_x).data[0];
        (mean, var)
    }
}

/// 特征:φ(x) = (1, x)——一元线性回归。
fn phi(x: f64) -> Matrix {
    Matrix::from_vec(vec![1.0, x], 2, 1)
}

fn mse(yhats: &[f64], ys: &[f64]) -> f64 {
    let n = ys.len();
    let se: f64 = yhats.iter().zip(ys.iter()).map(|(a, b)| (a - b).powi(2)).sum();
    se / n as f64
}

fn main() {
    let out_dir = "../../../frontend/public/images/series/rust-ml-07-bayesian";
    std::fs::create_dir_all(out_dir).expect("创建输出目录失败");

    // ---- 1. 数据:y = 2x + 1 + ε,ε ~ N(0, 0.3²),20 个点 ----
    let mut rng = XorShift::new(42);
    let n = 20usize;
    let mut xs = Vec::with_capacity(n);
    let mut ys = Vec::with_capacity(n);
    for _ in 0..n {
        let x = rng.next_f64() * 5.0;
        xs.push(x);
        ys.push(2.0 * x + 1.0 + 0.3 * rng.next_gauss());
    }
    let phi_all = Matrix::from_vec(
        xs.iter().flat_map(|x| [1.0, *x]).collect(),
        n,
        2,
    );

    let model = BayesReg::new(0.09, 1.0); // σ=0.3, τ=1(弱先验)
    println!("σ² = {:.2},τ² = {:.2},等价岭回归 λ = σ²/τ² = {:.4}", model.sigma2, model.tau2, model.ridge_lambda());

    // ---- 2. 序贯学习:数据一个个到来,后验如何收缩 ----
    println!("\n[N 个点时]  w0 均值±std        w1 均值±std         x*=2.5 处 ±2σ 带宽");
    for &k in &[0usize, 1, 5, 20] {
        let (mu, sigma) = if k == 0 {
            // N=0:后验 = 先验 N(0, τ²I)
            (Matrix::zeros(2, 1), Matrix::from_vec(vec![model.tau2, 0.0, 0.0, model.tau2], 2, 2))
        } else {
            let sub = Matrix::from_vec(
                xs[..k].iter().flat_map(|x| [1.0, *x]).collect(),
                k,
                2,
            );
            model.fit(&sub, &ys[..k])
        };
        let (_, var) = model.predict(&phi(2.5), &mu, &sigma);
        let s0 = sigma.get(0, 0).sqrt();
        let s1 = sigma.get(1, 1).sqrt();
        println!(
            "  N={:<2}      {:+.3} ± {:.3}   {:+.3} ± {:.3}   {:>7.3}",
            k,
            mu.get(0, 0),
            s0,
            mu.get(1, 0),
            s1,
            2.0 * var.sqrt()
        );
    }

    // ---- 3. MAP 线与后验预测带 ----
    let (mu, sigma) = model.fit(&phi_all, &ys);
    let grid: Vec<f64> = (0..=100).map(|k| k as f64 * 0.05).collect();
    let mut mean_c = Vec::new();
    let mut upper_c = Vec::new();
    let mut lower_c = Vec::new();
    let mut prior_upper = Vec::new();
    let mut prior_lower = Vec::new();
    for &x in &grid {
        let (m, v) = model.predict(&phi(x), &mu, &sigma);
        let s = v.sqrt();
        mean_c.push((x, m));
        upper_c.push((x, m + 2.0 * s));
        lower_c.push((x, m - 2.0 * s));
        // 先验预测带:均值 0,方差 τ²‖φ‖² + σ²
        let pv = model.sigma2 + model.tau2 * (1.0 + x * x);
        prior_upper.push((x, 2.0 * pv.sqrt()));
        prior_lower.push((x, -2.0 * pv.sqrt()));
    }

    let ymin = -3.0;
    let ymax = 14.0;
    let mut c = Canvas::new(560.0, 380.0, 0.0, 5.0, ymin, ymax);
    c.axes("x", "y");
    c.band(&prior_upper, &prior_lower, "#2563eb");
    c.band(&upper_c, &lower_c, PALETTE[0]);
    c.polyline(&mean_c, PALETTE[0], 2.2);
    c.dots(&xs.iter().copied().zip(ys.iter().copied()).collect::<Vec<_>>(), "#16161d", 3.5);
    c.legend(&[("prior band ±2σ", "#2563eb"), ("posterior band ±2σ", PALETTE[0]), ("MAP 线", PALETTE[0])]);
    let p1 = format!("{out_dir}/prior-vs-posterior.svg");
    c.save(&p1);

    // ---- 4. 不确定度随位置变化:数据密处窄,远处升高 ----
    let mut var_curve = Vec::new();
    for &x in &grid {
        let (_, v) = model.predict(&phi(x), &mu, &sigma);
        var_curve.push((x, v.sqrt()));
    }
    let mut c = Canvas::new(560.0, 320.0, 0.0, 5.0, 0.0, 0.0);
    c.axes("x", "预测 σ(x*)");
    c.polyline(&var_curve, PALETTE[1], 2.0);
    c.dots(&xs.iter().map(|x| (*x, 0.05)).collect::<Vec<_>>(), "#16161d", 2.5);
    let p2 = format!("{out_dir}/uncertainty-vs-x.svg");
    c.save(&p2);

    // ---- 5. 与岭回归的等价性验证 ----
    let xt = phi_all.transpose();
    let mut a = xt.matmul(&phi_all);
    for j in 0..2 {
        a.data[j * 2 + j] += model.ridge_lambda();
    }
    let ridge_w = a.inverse().unwrap().matmul(&xt.matmul(&Matrix::from_vec(ys.to_vec(), n, 1)));
    let diff = (ridge_w.get(0, 0) - mu.get(0, 0)).abs() + (ridge_w.get(1, 0) - mu.get(1, 0)).abs();
    println!("\nMAP 估计与岭回归(λ=σ²/τ²) 的 |Δw| 之和 = {:.2e}", diff);
    let val_mse = mse(&(0..n).map(|i| model.predict(&phi(xs[i]), &mu, &sigma).0).collect::<Vec<_>>(), &ys);
    println!("训练集 MSE(MAP 预测)= {:.4}", val_mse);
    println!("图已写入:{p1}、{p2}");
}

实现要点:

  • fit 直接翻译精度矩阵公式:Φ⊤Φ\Phi^{\top}\Phi 整体除以 σ2\sigma^2,对角线加 1/τ21/\tau^2,求逆得 Σ\Sigma,再乘 Φ⊤y/σ2\Phi^{\top}y/\sigma^2 得 μ\mu。没有梯度下降——共轭先验的世界里”训练”就是一次矩阵运算。
  • predict 返回 (均值, 方差) 元组,一行矩阵链 ϕ⊤Σϕ\phi^{\top}\Sigma\phi 拿到参数不确定度。
  • XorShift::next_gauss 用 Box-Muller 把均匀采样升级成标准正态:z=−2ln⁡u1cos⁡(2πu2)z = \sqrt{-2\ln u_1}\cos(2\pi u_2),零依赖得到造高斯噪声的能力。
  • 序贯学习表用 &[0usize, 1, 5, 20] 驱动,N=0 时跳过 fit、直接构造先验——后验=先验是贝叶斯更新的恒等起点。
// src/main.rs 段二:迷你 SVG 绘图器(与系列前篇同一份实现,新增 band 置信带)
use std::fmt::Write as _;

// ===================== 迷你 SVG 绘图器(与系列前篇一致) =====================

/// 系列默认配色:朱橙 / 青绿 / 蓝 / 琥珀(取自站点设计令牌)。
const PALETTE: [&str; 4] = ["#d6491f", "#0f766e", "#2563eb", "#d97706"];

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 {
    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)
    }

    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)
            );
        }
    }

    /// 置信带:upper/lower 两条曲线围成的区域,半透明填充。
    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 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
        );
    }

    /// 图例:右上角依次画「色线 + 文字」。
    fn legend(&mut self, entries: &[(&str, &str)]) {
        let sample_w = 22.0;
        let line_h = 18.0;
        let x0 = self.w - self.pad_r - 78.0;
        let x_text = x0 + 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 _ = 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 失败");
    }
}

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)
    }
}

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 box_muller_moments() {
        let mut rng = XorShift::new(42);
        let samples: Vec<f64> = (0..4000).map(|_| rng.next_gauss()).collect();
        let mean = samples.iter().sum::<f64>() / 4000.0;
        let var = samples.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / 4000.0;
        assert!(mean.abs() < 0.05, "均值应接近 0: {mean}");
        assert!((var - 1.0).abs() < 0.1, "方差应接近 1: {var}");
    }

    #[test]
    fn posterior_matches_ridge_map() {
        let mut rng = XorShift::new(7);
        let xs: Vec<f64> = (0..30).map(|_| rng.next_f64() * 5.0).collect();
        let ys: Vec<f64> = xs.iter().map(|x| 2.0 * x + 1.0 + 0.1 * rng.next_gauss()).collect();
        let phi_all = Matrix::from_vec(xs.iter().flat_map(|x| [1.0, *x]).collect(), 30, 2);
        let model = BayesReg::new(0.01, 1.0);
        let (mu, _) = model.fit(&phi_all, &ys);
        // 岭回归闭式解
        let xt = phi_all.transpose();
        let mut a = xt.matmul(&phi_all);
        for j in 0..2 {
            a.data[j * 2 + j] += model.ridge_lambda();
        }
        let ridge = a.inverse().unwrap().matmul(&xt.matmul(&Matrix::from_vec(ys.to_vec(), 30, 1)));
        assert!((ridge.get(0, 0) - mu.get(0, 0)).abs() < 1e-8);
        assert!((ridge.get(1, 0) - mu.get(1, 0)).abs() < 1e-8);
    }

    #[test]
    fn no_data_posterior_is_prior() {
        let model = BayesReg::new(0.09, 2.0);
        let mu = Matrix::zeros(2, 1);
        let sigma = Matrix::from_vec(vec![2.0, 0.0, 0.0, 2.0], 2, 2);
        let (_, var) = model.predict(&phi(1.0), &mu, &sigma);
        // 先验预测方差 = σ² + τ²‖φ‖² = 0.09 + 2×2
        assert!((var - (0.09 + 4.0)).abs() < 1e-9);
    }

    #[test]
    fn predictive_var_at_least_noise() {
        let mut rng = XorShift::new(3);
        let xs: Vec<f64> = (0..10).map(|_| rng.next_f64() * 5.0).collect();
        let ys: Vec<f64> = xs.iter().map(|x| x + 0.2 * rng.next_gauss()).collect();
        let phi_all = Matrix::from_vec(xs.iter().flat_map(|x| [1.0, *x]).collect(), 10, 2);
        let model = BayesReg::new(0.04, 1.0);
        let (mu, sigma) = model.fit(&phi_all, &ys);
        let (_, var) = model.predict(&phi(2.0), &mu, &sigma);
        assert!(var >= model.sigma2, "预测方差至少为噪声方差");
    }
}

Rust 语法角:关联函数与构造函数

本篇的 BayesReg::new(0.09, 1.0)、Matrix::zeros(2, 1)、XorShift::new(42) 用的是 Rust 的关联函数写法:在 impl 块里但第一个参数不是 self,用 :: 调用而不是 .。它就是别的语言里”静态方法/构造函数”的位置——Rust 的惯例是不写 new 特殊化,任何名字都可以(Matrix::zeros、Matrix::from_vec、Matrix::identity),返回 Self 即构造。方法(有 self)与关联函数(无 self)的区别只看第一个参数,这是 Rust 把”构造”统一进普通函数哲学的小例子。详见《Rust 程序设计语言》ch05-03(方法语法)。

运行结果

cargo test(4 个用例:Box-Muller 的均值/方差、后验均值与岭回归闭式解一致到 1e-8、无数据时后验=先验、预测方差不低于噪声地板)全部通过后,cargo run:

σ² = 0.09,τ² = 1.00,等价岭回归 λ = σ²/τ² = 0.0900

[N 个点时]  w0 均值±std        w1 均值±std         x*=2.5 处 ±2σ 带宽
  N=0       +0.000 ± 1.000   +0.000 ± 1.000     5.418
  N=1       +0.546 ± 0.287   +0.000 ± 1.000     5.069
  N=5       +0.747 ± 0.223   +2.053 ± 0.126     0.711
  N=20      +0.977 ± 0.138   +1.972 ± 0.050     0.615

MAP 估计与岭回归(λ=σ²/τ²) 的 |Δw| 之和 = 6.66e-15
训练集 MSE(MAP 预测)= 0.0903
图已写入:../../../frontend/public/images/series/rust-ml-07-bayesian/prior-vs-posterior.svg、../../../frontend/public/images/series/rust-ml-07-bayesian/uncertainty-vs-x.svg

两张图:

先验带与后验带:数据教会模型收缩

预测不确定度随位置变化:数据密处窄,两端张开

怎么读这些数字和图

  • 序贯表是贝叶斯学习的完整电影:N=0 时后验就是先验(±1.000);N=1 时截距先动(±0.287)而斜率纹丝不动(0.000 ± 1.000)——一个点确实告诉不了你斜率,模型诚实地保留了无知;N=5 时斜率冲到 2.05±0.13,N=20 收敛到 1.97±0.05,贴着真值 2。每个系数自带的 ±std 就是”确定性”的量化。
  • 带宽 5.418 → 0.615:x=2.5 处 ±2σ 带的宽度随数据收缩近 9 倍。注意它收敛到的不是 0——0.61 里大部分是噪声地板 2σ≈0.6,剩下的参数不确定度已经很小。可预测的噪声,不该被消掉。
  • MAP = 岭回归,误差 6.66e-15:浮点精度级别的零。06 篇的”λ 对应先验强度”在代码里被钉死成恒等式——两种世界观在同一个公式处汇合。
  • 训练 MSE 0.0903 ≈ σ²=0.09:模型把可学的都学了,剩下的就是噪声,和 05 篇”噪声地板”的结论遥相呼应。
  • 第二张图是不确定度的空间地图:黑色小点是训练数据的 x 坐标,σ(x*) 曲线在数据密集区下凹、两端翘起。同一个模型,在数据腹地自信、在外推边界谦虚——这是点估计模型永远给不出的信息。

优缺点与适用场景

(抄清单 1.9 原文)

  • 优点:天然的不确定性量化、小样本更稳健、防止过拟合有原理性解释(来自先验,而非工程补丁)。
  • 缺点:共轭先验是奢侈品——模型复杂一点就得请 MCMC 或变分推断出山(清单 1.9 的求解工具箱);先验超参数 σ²、τ² 本身需要估计。
  • 适用场景:数据量小、需要置信度输出的场景——A/B 测试、实验分析、任何”预测值要配误差棒”的业务。

小结

我们把”参数估计”升级成了”信念更新”:先验是信念的起点,每个数据点把后验收紧一格,预测自带误差棒,MAP 与岭回归精确等价——正则化找到了它的概率论解释。但还有一个局限:我们仍然假设关系是线性的(在特征空间里),不确定性只来自参数。如果让函数本身成为随机变量呢?下一篇回归⑥:高斯过程回归——用核函数对函数建模,小样本预测的王者。