第34章 线性回归
34.1 问题从哪来#
上一章的 KMeans 能把数据自动分成几组。分组之后,我们知道某个点被分到哪一组。
但有一类问题,KMeans 回答不了:
- 知道学生每周学习多少小时,想预测他的考试成绩。
- 知道房子的面积,想预测它的价格。
- 知道广告投入多少钱,想预测销售额。
这些问题不是在分组,而是在"给一个数字,算出另一个数字"。输入和输出都可以当作连续数值处理。分组算法把世界切成几块,但这些问题是想在数据里找到一条趋势线。
34.2 先看一个例子#
假设调查了 6 个学生的每周学习时长和考试成绩:
| 学生 | 每周学习时长 $x$ | 考试成绩 $y$ |
|---|---|---|
| A | 2 | 55 |
| B | 4 | 60 |
| C | 6 | 68 |
| D | 8 | 75 |
| E | 10 | 82 |
| F | 12 | 88 |
把这 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)$$有些误差是正的(点在线上方),有些是负的(点在线下方)。如果直接把误差加起来,正负会抵消,总和可能是零,但线并不好。
标准做法是把每个误差平方再加起来,得到误差平方和(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,提取 x 和 y。内存访问是顺序的,缓存友好。
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$ 在已有数据的范围内(比如 2~12 小时),预测通常比远离数据范围时更有参考价值。如果新 $x$ 远超范围(比如 50 小时),预测就是外推(extrapolation),结果可能完全不对——直线会无限延伸,但现实中成绩不会无限上涨。
还有一个检查办法:不要把所有数据都拿来拟合。可以留出几条数据,只用前面一部分求 $a$ 和 $b$,再用留下的数据检查预测误差。用于拟合的叫训练数据,用来检查预测效果的叫测试数据。这不是新的 C 语法,只是检查模型是否可靠的一种办法。
如果一条线在训练数据上误差很小,但换到测试数据就很差,它记住的是这批数据的细节,不是真正稳定的趋势。这种现象叫过拟合。这里的线性回归很简单,也仍然需要这个直觉:模型不是在背答案,而是在找能预测新数据的规律。
34.9 常见坑#
坑 1:数据点太少。 两个点能完美拟合一条线(SSE = 0),但这条线的预测可靠性很弱,它只是把两个点连起来。至少需要 5~6 个点才能看出趋势。
坑 2:把相关当因果。 学习时长和成绩高度相关,但"学得多"不一定是"成绩高"的原因。可能是成绩好的学生更有动力多学。线性回归找到的是相关关系,不是因果关系。
坑 3:忘记检查 denom。 如果所有 $x$ 相同,分母为零。上面的代码做了判断,但很多网上的实现没有。直接除以零会得到 inf 或 NaN。
坑 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") 打开文件。如果 x 和 y 用 double,就用 while (fscanf(fp, "%lf %lf", &x, &y) == 2) 逐行读取。记得读完 fclose(fp)。文件不存在时 fopen 返回 NULL,要先检查。
Q4:画拟合图。 写一个函数,用字符在终端里画出散点和拟合直线的近似图。横轴是 $x$,纵轴是 $y$,数据点用 *,拟合直线用 / 或 -。
提示:用一个二维字符数组 canvas[HEIGHT][WIDTH] 做画布,先全填空格。然后根据坐标缩放把每个数据点映射到画布格子里,改成 *。拟合直线部分:从 x_min 到 x_max 逐步算 y = a * x + b,在对应格子里画 / 或 -。最后逐行打印画布。
Q5:加入数据库。 把数据点存进前面章节的数据库结构,从数据库里读取点来做回归。这就是"数据库 + 简单机器学习"的雏形。
提示:复用前面章节的 struct DB 结构,数据点的 x 和 y 分别存成记录的两个字段。用循环遍历 db->rows,把数据读进 points 数组,再调用 fit。数据库里的记录可以来自文件、命令输入或之前章节的任何数据源。
34.11 全课程回顾#
从第 1 章到第 34 章,走过的路是这样的:
内存和数据。 第 1 章从内存格子开始,知道了数据在计算机里是一排字节,每个字节有地址,同一串字节按不同规则解释就是不同类型的数据。第 2 章引入变量和类型,第 3 章用判断和循环处理多次输入。
函数和结构。 第 4 章把逻辑拆成函数,第 5 章用数组管理一批成绩,第 6 章处理字符串,第 7 章用结构体表示一条学生记录,第 8 章把结构体数组组织成一张小表。
文件、调用栈、指针和动态容量。 第 9 章把结构体数组保存到文件,再从文件读回来。第 10 章用传地址操作外面的表,再用 malloc、realloc、free 让学生表的容量可以增长,也把指针和堆内存连到一起。
数据结构和算法。 第 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 自动分组。这是整个课程的收尾作品。