第34章 线性回归

34.1 问题从哪来#

上一章的 KMeans 能把数据自动分成几组。分组之后,我们知道某个点被分到哪一组。

但有一类问题,KMeans 回答不了:

  • 知道学生每周学习多少小时,想预测他的考试成绩。
  • 知道房子的面积,想预测它的价格。
  • 知道广告投入多少钱,想预测销售额。

这些问题不是在分组,而是在"给一个数字,算出另一个数字"。输入和输出都可以当作连续数值处理。分组算法把世界切成几块,但这些问题是想在数据里找到一条趋势线。


34.2 先看一个例子#

假设调查了 6 个学生的每周学习时长和考试成绩:

学生每周学习时长 $x$考试成绩 $y$
A255
B460
C668
D875
E1082
F1288

把这 6 个点画在坐标系里,横轴是学习时长,纵轴是成绩:

6 个数据点的散点图,横轴学习时长,纵轴成绩,点大致呈上升趋势

这些点不是随机散开的——学习时间越长,成绩越高。点大致沿着一条看不见的线排列。

如果能把这条线找出来,给一个新的学习时长,就能沿着这条线读出预测的成绩。比如学习 7 小时,预测成绩大约是多少?

这就是线性回归要做的事:给散点图找一条最佳拟合直线。


34.3 拟合是什么意思#

“找一条线"在数学上就是找两个数:斜率 $a$ 和截距 $b$,使得 $y = ax + b$ 尽量贴近所有数据点。

什么叫"尽量贴近”?对每个数据点,用线算出的 $\hat{y}_i = a x_i + b$ 和真实值 $y_i$ 之间有差距。这个差距叫误差(error):

$$e_i = y_i - (a x_i + b)$$

同一个 x 上真实值和预测值之间的纵向差值就是误差

有些误差是正的(点在线上方),有些是负的(点在线下方)。如果直接把误差加起来,正负会抵消,总和可能是零,但线并不好。

标准做法是把每个误差平方再加起来,得到误差平方和(sum of squared errors, SSE):

$$SSE = \sum_{i=1}^{n} (y_i - a x_i - b)^2$$

SSE 越小,说明这条线在所有数据点上的纵向误差平方和越小。“最佳拟合"就是让 SSE 最小的那条线。


34.4 怎么算最佳的 a 和 b#

对 $SSE$ 分别对 $a$ 和 $b$ 求偏导数,令它们等于零,可以推导出闭式解(不需要迭代,直接算):

$$a = \frac{n \sum x_i y_i - \sum x_i \sum y_i}{n \sum x_i^2 - (\sum x_i)^2}$$$$b = \frac{\sum y_i - a \sum x_i}{n}$$

这章直接使用这两个公式。你可以先观察它们需要哪些量:$\sum x_i$、$\sum y_i$、$\sum x_i y_i$、$\sum x_i^2$ 和数据点个数 $n$。程序遍历一遍数据,就能把这些值全部算出来。


34.5 最小实验#

下面这个小实验读入若干 $(x, y)$ 数据点,算出最佳拟合直线的 $a$ 和 $b$,然后用这条线预测新的 $x$ 对应的 $y$。重点看 fit 函数里那一轮求和。

#include <stdio.h>
#include <math.h>

#define MAX_POINTS 100
#define EPSILON 1e-12

struct Point {
    double x;
    double y;
};

struct LinearModel {
    double a;   // 斜率
    double b;   // 截距
};

int read_points(struct Point pts[], int max)
{
    int n = 0;
    printf("Enter number of data points: ");
    if (scanf("%d", &n) != 1) {                // 读取数据点个数
        printf("Invalid number format\n");
        return 0;
    }
    if (n > max) {                              // 检查是否超出最大容量
        printf("Max %d points allowed\n", max);
        return 0;
    }
    if (n < 2) {                                // 至少需要 2 个点才能拟合直线
        printf("Need at least 2 points\n");
        return 0;
    }
    for (int i = 0; i < n; i++) {               // 逐个读入每个数据点的 x 和 y
        printf("Point %d (x y): ", i + 1);
        if (scanf("%lf %lf", &pts[i].x, &pts[i].y) != 2) {
            printf("Invalid point format\n");
            return 0;
        }
    }
    return n;                                   // 返回成功读入的点数
}

struct LinearModel fit(const struct Point pts[], int n)
{
    double sum_x  = 0.0;                         // ∑x
    double sum_y  = 0.0;                         // ∑y
    double sum_xy = 0.0;                         // ∑xy
    double sum_x2 = 0.0;                         // ∑x²

    for (int i = 0; i < n; i++) {                // 一次遍历,同时累加四个统计量
        sum_x  += pts[i].x;
        sum_y  += pts[i].y;
        sum_xy += pts[i].x * pts[i].y;
        sum_x2 += pts[i].x * pts[i].x;
    }

    struct LinearModel m;
    double denom = n * sum_x2 - sum_x * sum_x;   // 分母 = n·∑x² - (∑x)²
    if (fabs(denom) < EPSILON) {
        // 所有 x 相同,无法拟合
        m.a = 0.0;                               // 斜率设为 0
        m.b = sum_y / n;                         // 截距取 y 的均值
    } else {
        m.a = (n * sum_xy - sum_x * sum_y) / denom;  // 斜率闭式解
        m.b = (sum_y - m.a * sum_x) / n;             // 截距闭式解
    }
    return m;
}

double predict(const struct LinearModel *m, double x)
{
    return m->a * x + m->b;                      // y = ax + b 预测值
}

void print_model(const struct LinearModel *m)
{
    printf("Fitted line: y = %.4f * x + %.4f\n", m->a, m->b);  // 输出模型方程
}

double compute_sse(const struct Point pts[], int n, const struct LinearModel *m)
{
    double sse = 0.0;                            // 误差平方和 SSE = ∑(yᵢ - ŷᵢ)²
    for (int i = 0; i < n; i++) {
        double err = pts[i].y - predict(m, pts[i].x);  // 误差 = 真实 y - 预测 y
        sse += err * err;                        // 累加误差平方
    }
    return sse;
}

void print_residuals(const struct Point pts[], int n, const struct LinearModel *m)
{
    printf("\n%-6s %8s %8s %8s\n", "Point", "True y", "Pred y", "Error");  // 表头
    printf("--------------------------------------\n");
    for (int i = 0; i < n; i++) {                // 逐点计算预测值并输出残差
        double y_hat = predict(m, pts[i].x);     // 预测值 ŷ = ax + b
        printf("%-6d %8.2f %8.2f %8.2f\n",
               i + 1, pts[i].y, y_hat, pts[i].y - y_hat);
    }
}

int main(void)
{
    struct Point pts[MAX_POINTS];
    int n = read_points(pts, MAX_POINTS);        // 1. 读入数据点
    if (n == 0) return 1;

    struct LinearModel m = fit(pts, n);          // 2. 拟合线性模型,计算 a 和 b
    print_model(&m);                             // 3. 输出模型方程

    printf("Sum of squared errors (SSE): %.4f\n", compute_sse(pts, n, &m));  // 4. 计算并输出 SSE

    print_residuals(pts, n, &m);                 // 5. 打印每个点的残差明细

    double new_x;
    printf("\nEnter x to predict: ");
    if (scanf("%lf", &new_x) != 1) {             // 6. 读取新 x,预测对应的 y
        printf("Invalid x format\n");
        return 1;
    }
    printf("Predicted y = %.4f\n", predict(&m, new_x));

    return 0;
}

34.6 编译运行#

保存为 linear_regression.c,编译:

$ gcc linear_regression.c -o linear_regression -lm
$ ./linear_regression

用前面的 6 个学生数据运行:

Enter number of data points:
$ 6
Point 1 (x y):
$ 2 55
Point 2 (x y):
$ 4 60
Point 3 (x y):
$ 6 68
Point 4 (x y):
$ 8 75
Point 5 (x y):
$ 10 82
Point 6 (x y):
$ 12 88
Fitted line: y = 3.4000 * x + 47.5333
Sum of squared errors (SSE): 2.1333

Point     True y   Pred y    Error
--------------------------------------
1        55.00    54.33      0.67
2        60.00    61.13     -1.13
3        68.00    67.93      0.07
4        75.00    74.73      0.27
5        82.00    81.53      0.47
6        88.00    88.33     -0.33

Enter x to predict:
$ 7
Predicted y = 71.3333

斜率 $a \approx 3.40$,意思是每周多学 1 小时,成绩大约提高 3.4 分。截距 $b \approx 47.53$,是 $x=0$ 时的基准成绩。用这条线预测学习 7 小时的成绩,得到约 71.33 分。

拟合直线穿过散点图,点大致分布在直线两侧


34.7 数据/内存/流程里发生了什么#

34.7.1 fit 函数做了什么#

fit 函数遍历一次所有数据点,累加四个值:

变量含义公式
sum_x所有 $x$ 之和$\sum x_i$
sum_y所有 $y$ 之和$\sum y_i$
sum_xy所有 $x_i y_i$ 之和$\sum x_i y_i$
sum_x2所有 $x_i^2$ 之和$\sum x_i^2$

这四个累加值代入公式,一步算出 $a$ 和 $b$。不需要循环迭代,不需要矩阵运算,时间复杂度是 $O(n)$。

34.7.2 内存布局#

6 个 Point 结构体在内存里连续排列。在常见 64 位环境里,一个 Point 通常占 16 字节(两个 double,各 8 字节);实际大小可以用 sizeof(struct Point) 验证。

地址         内容
0x1000      pts[0].x = 2.0     (8 字节)
0x1008      pts[0].y = 55.0    (8 字节)
0x1010      pts[1].x = 4.0     (8 字节)
0x1018      pts[1].y = 60.0    (8 字节)
...         ...

pts 指向数组开头,fit 函数通过下标从前往后访问每个 Point,提取 xy。内存访问是顺序的,缓存友好。

6 个 Point 结构体在内存中连续排列,fit 函数顺序遍历

34.7.3 为什么用 double 而不是 float#

成绩和学习时长用 float 也够,但 $x_i^2$ 和 $x_i y_i$ 的乘积会放大误差。6 个点时差别不大,数据点增多时,float 的 7 位有效数字可能不够。double 有 15~16 位有效数字,求和累积时更稳定。

34.7.4 denom 等于零的情况#

分母 $n \sum x_i^2 - (\sum x_i)^2$ 在所有 $x$ 值相同时等于零。这时候所有点排成一条竖线,任何斜率都无法拟合——$x$ 没有变化,就无法解释 $y$ 的变化。代码里用一个很小的 EPSILON 判断分母是否接近 0,退化成取 $y$ 的均值。


34.8 预测新点#

有了 $a$ 和 $b$,预测就是一行计算:$y = a \cdot x + b$。predict 函数做的就是这件事。

新输入的 x 沿拟合直线找到对应的 y,标出预测位置

预测的可靠性取决于数据点的分布。如果新 $x$ 在已有数据的范围内(比如 2~12 小时),预测通常比远离数据范围时更有参考价值。如果新 $x$ 远超范围(比如 50 小时),预测就是外推(extrapolation),结果可能完全不对——直线会无限延伸,但现实中成绩不会无限上涨。

还有一个检查办法:不要把所有数据都拿来拟合。可以留出几条数据,只用前面一部分求 $a$ 和 $b$,再用留下的数据检查预测误差。用于拟合的叫训练数据,用来检查预测效果的叫测试数据。这不是新的 C 语法,只是检查模型是否可靠的一种办法。

如果一条线在训练数据上误差很小,但换到测试数据就很差,它记住的是这批数据的细节,不是真正稳定的趋势。这种现象叫过拟合。这里的线性回归很简单,也仍然需要这个直觉:模型不是在背答案,而是在找能预测新数据的规律。


34.9 常见坑#

坑 1:数据点太少。 两个点能完美拟合一条线(SSE = 0),但这条线的预测可靠性很弱,它只是把两个点连起来。至少需要 5~6 个点才能看出趋势。

坑 2:把相关当因果。 学习时长和成绩高度相关,但"学得多"不一定是"成绩高"的原因。可能是成绩好的学生更有动力多学。线性回归找到的是相关关系,不是因果关系。

坑 3:忘记检查 denom。 如果所有 $x$ 相同,分母为零。上面的代码做了判断,但很多网上的实现没有。直接除以零会得到 infNaN

坑 4:用 float 导致精度不够。 数据点多、数值大时,float 的求和累积误差明显。涉及乘法再求和的场景,优先用 double

坑 5:外推太远。 拟合直线只在数据覆盖的范围内有意义。把直线延伸到很远的 $x$,预测值可能荒谬。比如学习 100 小时,按公式成绩会到 387.5 分左右。

坑 6:只看训练误差。 用来拟合的数据当然容易表现好。留出几条测试数据,才能看到这条线对新数据有没有预测能力。


34.10 自己试试看#

Q1:换个数据集。 用手动输入或 scanf 输入 10 组自己编的 $(x, y)$,观察 $a$ 和 $b$ 的变化。如果点很散(没有明显趋势),$a$ 会接近多少?

提示:故意输入一堆散乱的点,比如 (1, 5) (2, 9) (3, 2) (4, 8) (5, 3)。点之间没有明显趋势时,$a$ 会接近 0,$b$ 会接近所有 $y$ 的平均值——相当于拟合了一条水平线,说明 x 对预测 y 帮助不大。

Q2:计算 R²。 决定系数 $R^2$ 衡量拟合的好坏。公式是 $R^2 = 1 - \frac{SSE}{SST}$,其中 $SST = \sum (y_i - \bar{y})^2$。给 fit 函数加上 $R^2$ 的输出,$R^2$ 越接近 1,拟合越好。

提示:在 fit 函数里先遍历一遍数据算出 y_mean,再遍历一遍累加 SST。$R^2 = 0.9$ 说明 90% 的 y 变化可以被 x 解释;$R^2$ 接近 0 说明拟合不比直接猜平均值强多少。

Q3:文件输入。 把数据点存成文件,每行两个数字(x y),改 read_points 从文件读取而不是从键盘输入。

提示:用 FILE *fp = fopen("points.txt", "r") 打开文件。如果 xydouble,就用 while (fscanf(fp, "%lf %lf", &x, &y) == 2) 逐行读取。记得读完 fclose(fp)。文件不存在时 fopen 返回 NULL,要先检查。

Q4:画拟合图。 写一个函数,用字符在终端里画出散点和拟合直线的近似图。横轴是 $x$,纵轴是 $y$,数据点用 *,拟合直线用 /-

提示:用一个二维字符数组 canvas[HEIGHT][WIDTH] 做画布,先全填空格。然后根据坐标缩放把每个数据点映射到画布格子里,改成 *。拟合直线部分:从 x_minx_max 逐步算 y = a * x + b,在对应格子里画 /-。最后逐行打印画布。

Q5:加入数据库。 把数据点存进前面章节的数据库结构,从数据库里读取点来做回归。这就是"数据库 + 简单机器学习"的雏形。

提示:复用前面章节的 struct DB 结构,数据点的 xy 分别存成记录的两个字段。用循环遍历 db->rows,把数据读进 points 数组,再调用 fit。数据库里的记录可以来自文件、命令输入或之前章节的任何数据源。


34.11 全课程回顾#

从第 1 章到第 34 章,走过的路是这样的:

内存和数据。 第 1 章从内存格子开始,知道了数据在计算机里是一排字节,每个字节有地址,同一串字节按不同规则解释就是不同类型的数据。第 2 章引入变量和类型,第 3 章用判断和循环处理多次输入。

函数和结构。 第 4 章把逻辑拆成函数,第 5 章用数组管理一批成绩,第 6 章处理字符串,第 7 章用结构体表示一条学生记录,第 8 章把结构体数组组织成一张小表。

文件、调用栈、指针和动态容量。 第 9 章把结构体数组保存到文件,再从文件读回来。第 10 章用传地址操作外面的表,再用 mallocreallocfree 让学生表的容量可以增长,也把指针和堆内存连到一起。

数据结构和算法。 第 11 章用链表解决数组中间插删要搬数据的问题。第 12 章用栈实现撤销,第 13 章用队列表示先来先处理。第 14 章用两个栈处理四则表达式。第 15 章排序,第 16 章二分查找,第 17 章二叉搜索树,第 18 章哈希表,都是围绕"数据怎么放,查找和修改才更合适"这个问题展开。

程序检查和数据库生长。 第 19 章用操作计数、测试和不变量检查程序。第 20 章从数组版数据库开始,把插入、查找、删除、列出组织成一组函数。第 21 章改成动态数据库表,第 22 章换成链表版存储,第 23 章把外部调用稳定成接口。第 24 章把接口拆到头文件和实现文件里。第 25 章加有序索引,第 26 章把记录持久化到文件并重建索引,第 27 章用日志恢复崩溃前的操作,第 28 章让用户通过命令协议操作数据库,第 29 章用状态机解析带引号的参数。

从数据库到相似查询。 第 30 章把记录变成向量,第 31 章计算距离和相似度,第 32 章用 Top-K 和 KNN 找最近的几条记录。数据库不再只能回答"哪个 id 等于 1001”,也能回答"哪些记录离这个查询向量最近"。

简单机器学习。 第 33 章用 KMeans 按距离自动分组,第 34 章用线性回归拟合一条直线来预测连续数值。它们都没有离开前面那些基础:数组存数据,循环扫数据,函数封装计算,结构体组织记录。

可以把整套内容压成这张表:

阶段核心问题主要工具
基础数据数据怎么表示、怎么流动内存、变量、类型、判断、循环
批量数据一批记录怎么组织函数、数组、字符串、结构体、文件
动态存储容量和插删怎么处理指针、动态数组、链表
顺序和查找数据怎么放才查得快栈、队列、排序、二分、树、哈希
小数据库记录怎么管理、保存和恢复测试、不变量、DB 接口、模块、索引、持久化、日志、命令
相似和预测数据怎么比较、分组和预测向量、距离、Top-K、KMeans、线性回归

回头看,C 语言不是一堆孤立语法。int、数组、结构体、指针、文件、链表、排序、索引、日志、向量和误差,都在回答同一类问题:数据放在哪里,怎么找到它,怎么改变它,怎么比较它,怎么从它身上得出结论。


阶段项目#

全课程的最后一次整合练习是 阶段项目 5:相似记录查询工具:把记录变成向量,用欧氏距离找 Top-K 最相似的记录,还能用 KMeans 自动分组。这是整个课程的收尾作品。