用 Rust 从头实现机器学习算法·回归①:OLS 与梯度下降

53 分钟阅读 rust-ml-from-scratch · 3
Rust机器学习

用 Rust 从头实现机器学习算法·回归①:OLS 与梯度下降

基础篇备好了 Matrix,从本篇起进入系列的主线:逐个实现算法清单上的机器学习算法。第一个是最古老也最重要的——一元线性回归。整个过程分三步:先定义”什么叫拟合得好”,再算出”往哪个方向改参数会更好”,最后写一个循环反复改进。数学只需要链式法则,推导一步步展开。

核心思想

假设我们有一组数据点 (xi,yi)(x_i, y_i),i=1,…,ni = 1, \dots, n,大致落在一条直线附近。线性回归要找到这条直线的两个参数:斜率 ww 和截距 bb,使预测值

y^i=wxi+b\hat{y}_i = w x_i + b

尽可能接近真实的 yiy_i。找到之后,给定新的 xx 就能预测 yy——这就是”学习”。

算法清单对 OLS 的原文描述是:

核心思想:找到一条直线(或超平面),使所有样本点的「残差平方和」最小——最小二乘(Ordinary Least Squares)。

顺带说明一个背景:OLS 其实存在闭式解 β^=(X⊤X)−1X⊤y\hat{\boldsymbol{\beta}} = (X^\top X)^{-1} X^\top \boldsymbol{y}(正规方程,第二篇埋下过伏笔),但本篇刻意不用它,而是从损失函数出发用梯度下降迭代求解。原因是梯度下降是后面所有没有闭式解的模型(逻辑回归、神经网络……)的通用底座,值得先在只有 w,bw, b 两个参数的最简单地形上走一遍;“损失 → 梯度 → 更新”这三板斧将贯穿整个系列。

数学推导

第一步:损失函数——量化”拟合得好不好”

我们需要一个数来衡量”当前这条直线离数据有多远”。最直接的想法是把每个样本的误差加起来:∑(yi−y^i)\sum (y_i - \hat{y}_i)。但误差有正有负,会互相抵消——一条很离谱的直线,正负误差可能恰好抵消成零。

解决办法是平方:误差越大惩罚越重,正负号也消失了。

L(w,b)=12n∑i=1n(yi−(wxi+b))2L(w, b) = \frac{1}{2n} \sum_{i=1}^{n} \big( y_i - (w x_i + b) \big)^2

这个 LL 叫损失函数(均方误差的一种写法),它是参数 w,bw, b 的函数:每组 (w,b)(w, b) 都对应一个”分数”。前面的系数有两个小用意:除以 nn 让损失不随样本量膨胀;多除的 22 让下一步求导时系数恰好消掉。

把 LL 画出来,是一座抛物面,有且只有一个最低点。机器学习的”训练”,就是在这座山上从任意起点走到谷底。

第二步:求梯度——算出”往哪走”

站在山上某一点,想知道往哪个方向走能最快下山,需要偏导数:只把 ww(或 bb)当作变量,其余视为常数,看 LL 变化有多快。

对 ww 求偏导。外层是平方、内层是 wxi+bw x_i + b,用链式法则一层层剥:

∂L∂w=12n∑i=1n2(wxi+b−yi)⋅xi=1n∑i=1n(wxi+b−yi)xi\frac{\partial L}{\partial w} = \frac{1}{2n} \sum_{i=1}^{n} 2 \big( w x_i + b - y_i \big) \cdot x_i = \frac{1}{n} \sum_{i=1}^{n} \big( w x_i + b - y_i \big) x_i

对 bb 求偏导几乎一样,只是内层对 bb 的导数是 11:

∂L∂b=1n∑i=1n(wxi+b−yi)\frac{\partial L}{\partial b} = \frac{1}{n} \sum_{i=1}^{n} \big( w x_i + b - y_i \big)

读懂这两个式子的直觉:每个样本的”预测误差” wxi+b−yiw x_i + b - y_i 就是它对参数的推动力。误差的平均值(第二式)推着 bb 走;误差按 xix_i 加权后的平均值(第一式)推着 ww 走——xix_i 越大的样本对斜率的”发言权”越大。这与常识吻合:离原点远的点,对直线倾斜程度的影响更明显。

第三步:梯度下降——沿最陡方向下山

有了方向,更新规则只有一行:参数沿梯度的反方向走一小步,步长由学习率 η\eta 控制。

w←w−η∂L∂w,b←b−η∂L∂bw \leftarrow w - \eta \frac{\partial L}{\partial w}, \qquad b \leftarrow b - \eta \frac{\partial L}{\partial b}

为什么是减而不是加?梯度指向 LL 增长最快的方向,我们要下山,自然反着走。η\eta 太大可能一步跨过谷底来回震荡,太小则收敛缓慢——这是回归②要展开的话题。

Rust 实现

思路齐了,写代码。整个程序只依赖标准库,cargo run 即可运行。为了不引入外部依赖,我们手写一个经典的 xorshift 伪随机数生成器来造数据;文末的可视化则由一个内联的迷你 SVG 绘图器完成,两张图直接写进站点的静态资源目录。下面两段代码拼接起来就是完整的 src/main.rs。

首先是数据、训练与可视化主流程:

// 回归①:OLS 线性回归与梯度下降 —— 单文件实现(仅依赖 std)。
// 造数据 y = 3x + 2 + ε(xorshift64,种子 42),用梯度下降训练 w、b,
// 再用内联的迷你 SVG 绘图器输出拟合直线与 loss 收敛曲线。

use std::fmt::Write as _;

fn main() {
    // ---- 1. 生成模拟数据 y = 3x + 2 + ε ----
    let mut rng = XorShift::new(42);
    let n = 200usize;
    let mut xs = Vec::with_capacity(n);
    let mut ys = Vec::with_capacity(n);
    for _ in 0..n {
        let x = rng.next_f64() * 10.0 - 5.0;    // x ∈ [-5, 5)
        let eps = (rng.next_f64() - 0.5) * 2.0; // 噪声 ε ∈ [-1, 1)
        xs.push(x);
        ys.push(3.0 * x + 2.0 + eps);
    }

    // ---- 2. 初始化参数与超参数 ----
    let mut w = 0.0;
    let mut b = 0.0;
    let lr = 0.01;     // 学习率 η
    let epochs = 2000; // 迭代轮数

    // ---- 3. 训练循环:每轮先算损失与梯度,再沿梯度反方向更新参数 ----
    let mut history: Vec<(f64, f64)> = Vec::with_capacity(epochs);
    for epoch in 1..=epochs {
        let mut dw = 0.0;
        let mut db = 0.0;
        let mut loss = 0.0;
        for i in 0..n {
            let diff = w * xs[i] + b - ys[i]; // 预测误差 ŷ - y
            loss += diff * diff;
            dw += diff * xs[i];
            db += diff;
        }
        loss /= 2.0 * n as f64;
        dw /= n as f64;
        db /= n as f64;
        history.push((epoch as f64, loss));
        w -= lr * dw; // w ← w - η·∂L/∂w
        b -= lr * db; // b ← b - η·∂L/∂b
        if epoch == 1 || epoch % 200 == 0 {
            println!("epoch {:>4} | loss {:.6} | w {:.4} | b {:.4}", epoch, loss, w, b);
        }
    }
    println!("最终: w = {:.4}, b = {:.4}", w, b);

    // ---- 4. 可视化:拟合直线 + loss 收敛曲线 ----
    let out_dir = "../../../frontend/public/images/series/rust-ml-03-linear-regression";
    std::fs::create_dir_all(out_dir).expect("创建输出目录失败");

    // fit.svg:数据散点 + 训练后的拟合直线同框
    let line = vec![(-5.0, w * -5.0 + b), (5.0, w * 5.0 + b)];
    let pts: Vec<(f64, f64)> = xs.iter().copied().zip(ys.iter().copied()).collect();
    let (ymin, ymax) = y_range(&pts, &line);
    let mut fit = Canvas::new(560.0, 380.0, -5.0, 5.0, ymin, ymax);
    fit.axes("x", "y");
    fit.dots(&pts, "#d6491f", 3.0);
    fit.polyline(&line, "#0f766e", 2.0);
    fit.save(&format!("{out_dir}/fit.svg"));

    // loss.svg:loss 随 epoch 的收敛曲线
    let max_loss = history.iter().map(|p| p.1).fold(0.0, f64::max);
    let mut loss_canvas = Canvas::new(560.0, 320.0, 0.0, epochs as f64, 0.0, max_loss);
    loss_canvas.axes("epoch", "loss");
    loss_canvas.polyline(&history, "#d6491f", 2.0);
    loss_canvas.save(&format!("{out_dir}/loss.svg"));

    println!("SVG 已生成: {out_dir}/fit.svg, {out_dir}/loss.svg");
}

/// 数据点与拟合线端点的 y 值范围(绘图器自己再加 5% 边距)。
fn y_range(pts: &[(f64, f64)], line: &[(f64, f64)]) -> (f64, f64) {
    let mut ymin = f64::INFINITY;
    let mut ymax = f64::NEG_INFINITY;
    for (_, y) in pts.iter().chain(line.iter()) {
        ymin = ymin.min(*y);
        ymax = ymax.max(*y);
    }
    (ymin, ymax)
}

/// xorshift64 伪随机数生成器:无需外部依赖,结果可复现。
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
    }

    /// 返回 [0, 1) 区间的 f64。
    fn next_f64(&mut self) -> f64 {
        (self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
    }
}

实现要点:

  • 训练循环与三个公式逐行对应:diff 是预测误差 y^−y\hat{y} - y;dw、db 累加后除以 nn,正是 ∂L/∂w\partial L / \partial w 与 ∂L/∂b\partial L / \partial b;最后两行更新式与数学符号一一对应。对照检查一遍,代码里没有一处”魔法”。
  • loss、dw、db 都在参数更新之前、基于本轮起点计算——先看清方向再迈步,顺序不能颠倒。
  • 相比旧版训练循环,这里多了一个 history:每轮的 (epoch,loss)(\text{epoch}, \text{loss}) 存下来,既是 loss 收敛曲线的数据源,也让第 1 轮与最后几轮的 loss 直接出现在文本输出里。
  • 可视化输出用相对路径加 create_dir_all:从 series/rust-ml/rust-ml-03-linear-regression/ 运行 cargo run,两张 SVG 直接写进站点的静态资源目录 frontend/public/images/series/rust-ml-03-linear-regression/。

然后是迷你 SVG 绘图器(与系列其他文章共用同一份实现,series/rust-ml/_plotter-demo/ 里也有它):


// ---- 迷你 SVG 绘图器:零依赖(仅 std),把数据写成 SVG 文件 ----
// 这是系列文章的共享绘图工具,每篇文章的代码中都会内联同一份实现。

/// 一张图:坐标映射 + 已累积的 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
        );
    }

    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
}

实现要点:

  • Canvas 持有数据坐标范围和一个不断累积 SVG 元素的 body 字符串;px / py 做数据坐标到像素的线性映射,py 里有一次翻转——屏幕的 y 轴向下,数学的 y 轴向上。
  • polyline、dots、axes、save 四个方法覆盖折线、散点、坐标轴与落盘;刻度步长取 1/2/2.5/5×10k1 / 2 / 2.5 / 5 \times 10^k 里”最好看”的一档。
  • 注意方法签名上的 &self 与 &mut self 之分:只读的取 &self,会往 body 里追加元素的取 &mut self——这不是风格偏好,而是 Rust 借用规则的直接体现,马上细说。

Rust 语法角:借用与可变借用

Python 里变量是指向对象的引用(名字),list 天生可变:几个变量可以指向同一个列表,通过任何一个都能改它,语言并不阻止你一边遍历、一边在别处修改——这类”别名 bug”往往要到运行时才暴露。

Rust 把这件事管了起来,规则只有两条:

  • &T 不可变借用:可以有很多个,但谁都不得修改原数据;
  • &mut T 可变借用:同一时刻只允许存在一个,且它活着的时候不允许有任何 &T 指向同一数据。

全部由编译器在编译期静态检查——违反即报错,而不是运行到那一行才崩。

回到训练循环,这两条规则体现得淋漓尽致:xs、ys 建好之后就再也不改,循环里读取 xs[i] 只是不可变借用,所以它们不需要 mut,&pts 还能同时借给多个绘图调用;而 w -= lr * dw 要修改 w,w 就必须声明为 mut——它拿到的是”独占的修改权”。绘图器的 API 是同一原则的对照组:dots 和 save 只读,签名是 &self;polyline 和 axes 会累积元素,签名是 &mut self。这套设计让数据竞争类错误在编译期就无处遁形。想深入可读《Rust 程序设计语言》第 4 章 References and Borrowing。

运行结果

运行 cargo run(随机数种子固定为 42,每次输出逐字节一致):

epoch    1 | loss 32.326143 | w 0.2093 | b 0.0061
epoch  200 | loss 0.202033 | w 2.9938 | b 1.6762
epoch  400 | loss 0.153697 | w 3.0136 | b 1.9466
epoch  600 | loss 0.152703 | w 3.0164 | b 1.9854
epoch  800 | loss 0.152683 | w 3.0168 | b 1.9910
epoch 1000 | loss 0.152682 | w 3.0169 | b 1.9918
epoch 1200 | loss 0.152682 | w 3.0169 | b 1.9919
epoch 1400 | loss 0.152682 | w 3.0169 | b 1.9919
epoch 1600 | loss 0.152682 | w 3.0169 | b 1.9919
epoch 1800 | loss 0.152682 | w 3.0169 | b 1.9919
epoch 2000 | loss 0.152682 | w 3.0169 | b 1.9919
最终: w = 3.0169, b = 1.9919
SVG 已生成: ../../../frontend/public/images/series/rust-ml-03-linear-regression/fit.svg, ../../../frontend/public/images/series/rust-ml-03-linear-regression/loss.svg

拟合结果(朱橙为数据点,青色为训练后的直线 y^=3.0169 x+1.9919\hat{y} = 3.0169\,x + 1.9919):

loss 收敛曲线(2000 轮全部记录):

怎么读这组数字与这两张图

  • loss 先快后慢:从 32.3332.33 一路降到 0.200.20 只用了一两百轮,之后几乎原地踏步。这是梯度下降的典型轨迹——离谷底越远梯度越大、步子越猛;接近谷底后梯度趋近于零,参数只做”微调”。loss.svg 里那条陡降后走平的曲线,就是”在抛物面上大步流星、然后小步挪向谷底”的足迹。
  • loss 停在 0.15 附近而不是 0:数据里有人工加入的噪声 ε\varepsilon,它本质上不可预测。0.150.15 附近就是噪声的”地板”:模型学会了规律(那条直线),但不可能把噪声也预测出来——做不到,也不该做到。
  • w→3.017w \to 3.017、b→1.992b \to 1.992:非常接近造数据时用的真值 w=3w = 3、b=2b = 2。不完全相等正是噪声使然——如果完美等于 3 和 2,反而要怀疑它把噪声也背下来了(过拟合的雏形)。
  • fit.svg 怎么看:青色直线斜穿朱橙散点带的中心,残差正负散布在直线两侧、没有系统性偏向哪一边——这正是”残差平方和最小”的肉眼版。可以顺手验证:直线在 x=0x = 0 处的高度约是 b≈2b \approx 2,xx 每加 1,yy 约加 3。

另外注意,本篇没有用基础篇的 Matrix:一元回归只有一个特征,用 Vec<f64> 足矣,代码更短、更聚焦梯度下降本身。等进入多元回归,把标量 xix_i 换成特征向量、把 Vec 换成 Matrix,梯度公式几乎原封不动——这正是基础篇末尾 y^=Xw\hat{\boldsymbol{y}} = X \boldsymbol{w} 的威力。

优缺点与适用场景

照抄算法清单 1.1 的原文评价:

  • 优点:可解释性强、有完整统计推断、闭式解无需迭代。
  • 缺点:对多重共线性和异常值敏感,无法表达非线性关系。
  • 适用场景:基线模型、需要解释系数的场景(经济学、社会科学)。

对照本篇的实验可以直观理解第一条:OLS 的每个系数都有明确含义(斜率 = xx 每加一单位 yy 的变化量),而且像本篇这样有闭式解的模型甚至可以一步到位、无需训练循环;它的缺陷则要等到多元回归与多项式回归篇才会真正现身。

小结

我们走完了机器学习的完整闭环:定义损失(最小二乘)→ 求梯度(链式法则)→ 迭代优化(梯度下降)→ 在模拟数据上验证收敛,并且第一次让结果可视化。全部代码只依赖标准库。下一篇进入回归②:把标量 xx 扩展成特征向量实现多元回归,直面两个真实问题——特征量纲悬殊时梯度下降为什么会来回震荡(特征缩放),以及学习率与正则化如何联手让模型既收敛得快又不死记硬背。