news 2026/9/8 1:34:20

线性回归从零实现:损失函数、梯度下降与模型评估全解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
线性回归从零实现:损失函数、梯度下降与模型评估全解析

学机器学习的人,十个里有八个是从线性回归开始的。这句话放在哪一年的课程目录里都成立。但真正有意思的是,十个里可能只有两三个在学完之后,能拍着胸脯说自己真的理解了正在做的事情。

我之前看过不少机器学习入门课程,有一类特别负责,会把线性回归拆成“代码 + 直觉”两条线来讲。有的课程用英语,有的用印地语,配上中文字幕。语言不一定听得懂,但你会在某个节点突然意识到:线性回归根本不是教你怎么画一条穿过散点的直线,而是在教一种建模的思维顺序——先想清楚要解决什么问题,再用数学语言定义什么叫做“差”,然后通过优化把这个“差”压到最小,最后还要检查这个“最小”到底够不够好。

这篇文章想复现的就是这种讲法。不绕开数学,但尽量用人话解释;不跳过代码,但保证你看完知道自己为什么这么写。先说结论:简单线性回归最值得你带走的,不是某个公式,也不是一段能跑的代码,而是“定义问题 → 构造损失 → 优化参数 → 评估结果”这条建模链路。后面学逻辑回归、决策树、神经网络,骨架还是它。

1. 先搞清楚线性回归解决的是哪一类问题

1.1 它不是“画一条直线”这么简单

很多人第一次看到线性回归的配图,心里想的是:这不就是在一堆散点里画一条离所有点都最近的直线吗?这个直觉没有错,但它漏掉了最重要的部分——这条直线不是凭空画出来的,它是在一个明确的优化目标下被“计算”出来的。

线性回归真正解决的问题是:给定一个输入特征 X,预测一个连续输出 y。典型的例子有:

  • 根据学习时长预测考试成绩;
  • 根据房屋面积预测房价;
  • 根据广告投放金额预测销售额;
  • 根据温度预测冰淇淋销量。

这些问题的共同点是:X 和 y 之间存在某种趋势,但趋势被噪声掩盖。同样复习五个小时,A 考了 90 分,B 可能只考了 78 分。如果我们坚持用一条直线去描述“复习时长和成绩”的关系,这条线不可能穿过每一个点,它只能尽量靠近所有点。

“尽量靠近”不是一句模糊的话,它必须被翻译成数学语言。翻译的结果就是损失函数。后面你会发现,机器学习的很多所谓“高级算法”,本质上也都是在做同一件事:把模糊的目标翻译成可计算、可优化的函数。

1.2 模型表达式和它背后的假设

简单线性回归的模型是:

y = w * X + b

w 是权重(斜率),表示 X 每变化一个单位,y 平均变化多少;b 是偏置(截距),表示 X 为 0 时的基准值。在 sklearn 中,模型训练完成后分别存为coef_intercept_

这个简单的式子背后有几个假设:

  1. 线性关系:X 和 y 的关系可以用直线近似。
  2. 误差独立同分布:每个样本的噪声互相独立,且拥有相同的分布。
  3. 误差均值为 0:模型在平均意义上没有系统性的高估或低估。

新手最容易忽略的是第一条。拿到数据的第一件事,应该是画散点图,用眼睛确认关系是不是近似线性的。如果数据呈现明显的曲线趋势,比如 y 随 X 增长得越来越快,直接用线性回归去拟合,结果会很差。这时候要么做特征变换,要么换模型,而不是硬调参数。

这里要澄清一个常见误解:线性回归里的“线性”是指在参数 w、b 上线性,不是指 X 本身一定要是直线关系。你可以把 X² 变成一个新特征再放进模型,模型照样能表达曲线。但在简单线性回归里,通常先不加这些花活,把最朴素的场景理解透再说。

2. 为什么误差要平方,参数怎么求

2.1 损失函数的三层含义

给定一组数据,理论上可以画出无数条直线。哪条最好?需要一个标准来回答。这个标准就是损失函数。

最常见的损失函数是均方误差(MSE):

loss = np.mean((w * X + b - y) ** 2)

为什么用平方而不是绝对值?至少有三个原因。

第一,平方误差对大误差更敏感。预测偏差 2 个单位和 4 个单位的损失分别是 4 和 16,而不是 2 和 4。这意味着模型优化时会优先去修正那些偏差很大的样本。这个性质有好有坏,好的是收敛方向通常更准确,坏的是模型会被极端异常值带偏。后面讲实战坑点时会再提到。

第二,平方误差处处可导。梯度下降的每一次更新都依赖于求导数。绝对值函数在 0 点不可导,虽然也可以凑合处理,但平方误差的导数形式干净,就是2 * (pred - y),实现起来非常方便。

第三,它和正态分布有天然联系。在误差服从正态分布的假设下,最小化平方误差等价于极大似然估计。这个理解在推导层面很有用,但初学阶段记住前两条就够了。

2.2 正规方程和梯度下降:两种求解路线

求 w 和 b 使损失最小,有两条典型路线。

正规方程(闭式解):直接对损失函数求导,令导数等于 0,解出 w 和 b。数学上一步到位,不需要迭代。特征很少、数据量不大时,计算非常快。热词里提到的“线性回归的正规方程解”,指的就是这个。

梯度下降(迭代解):先随便给 w 和 b 一个初始值,然后计算损失对 w 和 b 的偏导数,沿着导数下降的方向更新参数,重复很多次,直到损失收敛。

两者的关系,就像一个是直接算出目的地坐标,另一个是“看着地图一步步走”。特征少、数据量小时用正规方程很舒服;特征多、数据量大时,正规方程涉及矩阵求逆,计算代价高,梯度下降更可行。sklearn 的LinearRegression默认走的是最小二乘路线,内部用 SVD 实现,稳定性很好。

2.3 学习率:你迟早要面对的一个参数

梯度下降里最重要的超参数是学习率lr。它决定每一步往梯度反方向走多远。

  • 学习率太大,参数会在最优点附近来回震荡,损失曲线上下跳动,甚至发散。
  • 学习率太小,收敛速度慢到让人失去耐心,可能迭代几千轮还停在原地。

如果你拿到一份代码,运行后发现 loss 曲线在上下乱跳,优先怀疑学习率,而不是模型的公式写错了。一般从 0.01 到 0.001 这个范围开始试。不同数据集的量级不一样,学习率不是一个能一劳永逸写死的值。换一份数据,同样设 0.01,可能效果就完全不一样了。

3. 用 Python 从零实现并验证

3.1 生成一份带噪声的演示数据

为了不依赖真实数据集,先用 numpy 构造一份符合线性关系、但带着噪声的数据:

import numpy as np np.random.seed(42) X = np.linspace(0, 10, 50) true_w = 2.5 true_b = 1.0 y = true_w * X + true_b + np.random.normal(0, 1.5, size=50)

这份数据的真实关系是y = 2.5X + 1,我们加了一个标准差为 1.5 的噪声。噪声让问题变得真实——如果数据完全无噪声,两个点就能确定一条直线,也就无所谓“学习”了。

也可以顺手把散点图画出来,确认数据确实呈现一条带毛边的直线带:

import matplotlib.pyplot as plt plt.scatter(X, y, alpha=0.7) plt.xlabel("X") plt.ylabel("y") plt.show()

3.2 手写梯度下降:二十行代码理解核心

不调用任何机器学习库,用纯 numpy 实现一遍:

def compute_loss(X, y, w, b): pred = w * X + b return np.mean((pred - y) ** 2) def gradient_descent(X, y, w=0.0, b=0.0, lr=0.02, epochs=300): n = len(y) loss_history = [] for _ in range(epochs): pred = w * X + b dw = (2 / n) * np.dot(X, pred - y) db = (2 / n) * np.sum(pred - y) w -= lr * dw b -= lr * db loss_history.append(compute_loss(X, y, w, b)) return w, b, loss_history w, b, history = gradient_descent(X, y) print(w, b)

这里最需要看懂的是两个梯度的计算:

  • dw是损失对 w 的偏导,用 X 和误差向量做点积,再乘2 / n
  • db是损失对 b 的偏导,误差向量直接求和,再乘2 / n

误差向量pred - y是每次更新时最核心的信息:它告诉模型“你每个点偏了多少、偏的方向是什么”。梯度下降本质上就是反复利用这个误差修正参数。

跑完之后,w 和 b 应该落在接近 2.5 和 1.0 的位置。具体数值会因为迭代次数和噪声略有偏差,但你会观察到:随着训练轮数增加,loss 先快速下降,然后趋于平缓。这就是收敛。

history画出来,是确认训练过程是否健康的第一步:

plt.plot(history) plt.xlabel("epoch") plt.ylabel("loss") plt.show()

不要急着调参数。第一次跑通代码时,先画 loss 曲线,确认它是平滑下降的。如果曲线在跳,再去看学习率。

3.3 用 sklearn 对比验证

手写版本是为了训练直觉,工程上直接用 sklearn:

from sklearn.linear_model import LinearRegression model = LinearRegression() model.fit(X.reshape(-1, 1), y) print("w:", model.coef_[0]) print("b:", model.intercept_)

sklearn 的LinearRegression默认使用最小二乘法,不需要设置学习率和迭代次数。输出和手写梯度下降收敛后的结果基本一致。

这里有一个新手必踩的坑:直接把一维数组 X 传给fit(),结果报错。线性回归要求输入是二维的,特征必须按行排列。所以要用X.reshape(-1, 1),把一个长度为 50 的一维数组变成 50 行 1 列的矩阵。

如果你是在头歌这类实训平台上做线性回归作业,经常会看到“python 和 sklearn 混合版”的要求。思路和这里完全一样:先用 numpy 手写核心流程理解原理,再用 sklearn 验证结果。平台判分时既看重原理实现,也看重最终模型效果,所以两部分代码都要能跑通。遇到这类作业,与其去网上找现成答案,不如先把手写梯度下降和 sklearn 的对比逻辑彻底搞懂。

3.4 读懂参数含义,而不只是打印出来

拿到coef_intercept_之后,很多人就结束了。其实更值得做的一步是解读参数的业务含义:在这个例子里,w 约等于 2.5,意思是“当 X 每增加 1,y 平均增加 2.5”;b 约等于 1,意思是“当 X 等于 0 时,y 的基准值大约是 1”。

注意“平均”两个字。因为数据带噪声,参数不是真实的因果关系保证,它只是对当前数据最优的估计。能不能把 w 解释成因果效应,取决于你的数据和实验设计,这是另一个话题。但至少,你应该养成习惯:每次建模之后,用一句自然语言描述出模型的含义。能说清楚模型在算什么,才算真正理解了这个模型。

4. 评估线性回归:三个指标加一张图

4.1 R²、RMSE、MAE 各自说明什么

模型训练完,必须回答“到底好不好”的问题。常用的三个指标一起看:

import numpy as np from sklearn.metrics import r2_score, mean_squared_error, mean_absolute_error y_pred = model.predict(X.reshape(-1, 1)) r2 = r2_score(y, y_pred) rmse = np.sqrt(mean_squared_error(y, y_pred)) mae = mean_absolute_error(y, y_pred) print("R²:", r2) print("RMSE:", rmse) print("MAE:", mae)
指标它衡量什么特点
模型解释了 y 方差的百分之多少越接近 1 越好;为负数说明比直接猜均值还差
RMSE预测误差的标准差量级单位与 y 一致;对大误差敏感
MAE预测误差的平均绝对偏差单位与 y 一致;对异常值不那么敏感

三个指标不应该选一个用,而应该一起读。R² 看整体拟合度,RMSE 看“大错误有多大”,MAE 看“平均错误有多大”。如果 RMSE 明显大于 MAE,说明误差主要由少数极端样本贡献,这时候再回头检查异常值,会比继续调参更有价值。

4.2 残差图能告诉你指标看不见的问题

残差是y - y_pred,也就是真实值和预测值的差。画残差图时,横轴是预测值,纵轴是残差:

plt.scatter(y_pred, y - y_pred, alpha=0.7) plt.axhline(y=0, color="red", linestyle="--") plt.xlabel("predicted") plt.ylabel("residual") plt.show()

一个合格的线性回归模型,残差图应该是“随机散乱地分布在水平线 0 附近”,没有明显形状。

如果你看到残差呈现喇叭形——预测值越大,残差越分散——说明存在异方差性;如果残差呈现弯曲的弧线,说明你漏掉了非线性关系;如果有几个点明显远离整体,说明存在异常值。

很多人只看 R²,不看残差图。R² 很高时,残差图也可能暴露出模型在局部区域表现很差。建议把残差图固定成每次建模的必查项,这比记一百个评估指标都管用。

4.3 线性回归适合什么,不适合什么

线性回归适合:

  • 学习阶段理解建模流程,它是最干净的起点;
  • 特征和输出关系近似线性的回归问题;
  • 需要模型可解释性的业务场景,比如定价、风控、销售预测;
  • 作为更复杂模型的 baseline,先用简单模型把效果基准打出来。

不适合:

  • 高维稀疏特征,比如大规模文本向量;
  • 特征之间存在严重多重共线性的多元场景;
  • 输出是分类标签,那是逻辑回归的任务;
  • 数据量极大、特征极多且关系复杂,此时树模型或神经网络通常更合适。

线性回归在生产环境里经常不是最终模型,但它几乎永远是第一个模型。用简单的模型先建立 baseline,再用复杂模型对比提升,这是一个非常实用的工程习惯。

5. 实战中最容易翻车的五个环节

5.1 异常值把回归线拉偏

一个极端异常值就能把回归线“拽”偏,尤其是 x 方向离群的点,对斜率的影响极大。原因是平方误差对大误差的惩罚太重,异常值的残差在损失里占了很大权重,模型会把大量精力用来迁就它。

建模前先画散点图,用 IQR 或 z-score 做一次异常值筛查。但注意:是“排查”,不是“一律删除”。异常值有时携带重要信息,比如某个销售异常波动背后可能是活动或事故。先把异常值找出来,判断它是噪声还是信号,再做后续处理。

5.2 train_test_split 用错位置

如果直接把全部数据拿去fit,再拿同一批数据评估 R²,得到的是“在训练集上的表现”,通常虚高。正确流程是先用train_test_split划分数据,在训练集上fit,在测试集上评估:

from sklearn.model_selection import train_test_split X_train, X_test, y_train, y_test = train_test_split( X.reshape(-1, 1), y, test_size=0.2, random_state=42 ) model = LinearRegression() model.fit(X_train, y_train) y_test_pred = model.predict(X_test)

很多刚入门的人会问:为什么测试集上的分数总比训练集低?这是正常的。模型在训练集上有一定的“记忆”,只有在新数据上才能看出它的真实泛化能力。

5.3 数据泄露

数据泄露是一个更隐蔽的问题。如果你在划分训练测试集之前,先用全量数据计算均值、方差,做了归一化或填补缺失值,测试集的信息就已经渗进了训练过程。预处理器只能在训练集上fit,然后在测试集上transform

在简单线性回归里,数据泄露通常发生在归一化环节。虽然 sklearn 的LinearRegression自己不做归一化,但如果你在流水线里加了StandardScaler,就要严格按照“先划分,再 fit 预处理器”的顺序来。这一步错了,指标再好看都不能说明模型真的有效。

5.4 时间序列数据当独立样本

如果数据是时间序列,比如按天记录的销售额、温度、股票指标,随机划分训练测试集就会泄露未来信息。训练集里的某条数据可能出现在测试集时间范围之后,模型等于“偷看”了未来。

这时的正确做法是按时间顺序划分,不能用随机打乱。线性回归教程里很少有人提这件事,但它是真实业务里最常见的错误之一。看到数据里带日期,先停下来确认划分方式,再跑模型。

5.5 指标漂亮但不符合业务直觉

最后这个算不上技术问题,但非常常见。模型指标很好,参数方向却不符合业务常识,比如广告投放越多销售额反而下降,或者价格越高销量越高。这时候不要急着上线,先回去检查数据口径、单位、异常值,再看特征是否合理。

指标只是工具,业务判断才是终点。一个方向都错了的模型,指标再高也不能投入使用。建模不是跑完代码就结束,模型的输出必须能回到业务里接受解释和挑战。

5.6 一个可复用的排查链路

遇到线性回归结果不对时,按这个顺序排查,比直接瞎调参数可靠得多:

  1. 看现象:loss 不下降、loss 震荡、训练 R² 很高但测试很低、参数符号反常。
  2. 查输入:X 有没有reshape,有没有 NaN,数据是否按时间顺序却做了随机划分。
  3. 查预处理:异常值有没有处理,归一化是否在train_test_split之后。
  4. 查训练参数:学习率、迭代轮数是否合理。
  5. 查模型假设:画 y 对 X 的散点图确认线性性,画残差图确认噪声形态。

这个链路覆盖了从数据到实现到假设的完整链条。大多数“看起来很奇怪”的线性回归结果,根源都不是代码 bug,而是流程顺序或数据问题。

回到开头那句话:十个学机器学习的人里有八个从线性回归开始。这篇文章真正想留下的,不是代码,也不是数学推导,而是一套思维顺序:先定义问题,再构造损失,然后优化参数,最后评估模型。

这套顺序在逻辑回归里一样,在神经网络里也一样。差别只是损失函数变复杂了、参数变多了、求解方法更高级了,骨架没有变。如果你能从一段二十行的梯度下降代码里真切感受到“误差如何一步步把参数推到最优位置”,那线性回归这一课就算真正学完了。

下一步,建议你换一份真实数据,比如房价或销售数据,把上面的流程完整跑一遍。先别急着学下一个算法,把一个流程走到能独立解读结果、能自己排查错误,比囫囵吞枣地刷十个算法有用得多。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/8 1:33:53

工具站SEO监测指标体系:跳出率、转化率与Core Web Vitals实战指南

手里攥着一个工具站的流量后台,最折磨人的通常不是没流量,而是流量来了之后你分不清它到底算好还是算坏。同样一个跳出率,内容站看了要连夜改稿,工具站看了可能只是用户完事就走。我自己做过在线转换、生成器和数据辅助类工具&…

作者头像 李华
网站建设 2026/9/8 1:32:55

nRF52832 GCC编译实战:从搭建环境到烧录调试

简介:面向Nordic 52832低功耗蓝牙SoC开发者的GCC编译环境搭建资料合集,涵盖从工具链安装、环境变量配置到固件编译下载的完整流程,特别适合需要在Windows/Linux下使用开源工具链开发BLE物联网设备的嵌入式工程师。包体共27037个文件&#xff…

作者头像 李华
网站建设 2026/9/8 1:31:11

移动式车辆轴重检测仪:铝合金秤台在公路港口检测实战解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/8 1:31:05

daq-2.0.7.tar.gz 编译安装全攻略:从解压到 Snort 2.9 对接

简介:压缩包 daq-2.0.7.tar.gz 是 Snort 入侵检测系统数据采集组件 DAQ 的 2.0.7 版本源码包,面向网络安全管理、Snort 部署与二次开发人员,用于解决多样网络环境下数据包统一接入与获取的问题。包内共 74 个文件,23 个 C 头文件和…

作者头像 李华
网站建设 2026/9/8 1:31:02

Qt5实战:从零开发十字路口信号灯模拟器

简介:面向Linux环境下QT5初学者的十字路口红绿灯模拟程序,适用于嵌入式Linux开发、智能交通课程设计和交通信号控制实验等场景。程序基于QT5的图形界面与信号槽机制搭建,使用户能够通过自定义协议控制红绿灯各灯状态的切换,界面可…

作者头像 李华
网站建设 2026/9/8 1:29:35

GPS+INS组合导航:从卡尔曼滤波到工程落地

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华