机器学习:线性回归学习

本内容为个人学习心得整理,基于陈老师《机器学习》课程。

本篇内容理论较多。

1. 问题引入:从一个发硬币的小游戏说起

我们先来看一个简单的小游戏:老师每天给你一定数量的硬币,以下是过去四天的记录:

|----|------|------|------|------|------|
| 日期 | Day1 | Day2 | Day3 | Day4 | Day5 |
| 收入 | 1 | 2 | 3 | 4 | ? |

那么第五天会给多少个硬币呢?我们很可能会猜测是5枚

我们是依据什么来猜测的呢?不管是否意识到,过去四天的数据在我们脑海中已经形成了一个规律------每天比前一天多1枚。把这些数据画在坐标系中,就是一条完美的直线。

现在我们把游戏难度升级,数据变成这样:

|----|------|------|------|------|------|
| 日期 | Day1 | Day2 | Day3 | Day4 | Day5 |
| 收入 | 2 | 5 | 6 | 9 | ? |

这次规律就不那么明显了。把数据点画在坐标系中,我们发现这些点并不在一条直线上,但似乎又存在某种趋势。

这时候问题就来了:我们能不能找到一条直线,尽可能地 "穿过" 这些点,从而预测第五天的硬币数?

这就是线性回归要解决的核心问题。

2. 线性回归的定义

2.1 什么是线性回归

线性回归(Linear Regression) 是利用回归方程(函数)对一个或多个自变量(特征值)和因变量(目标值)之间关系进行建模的一种分析方式。

简单来说,线性回归就是要找到一个线性函数,用它来描述特征值和目标值之间的关系,从而对新的未知数据做出预测。

2.2 单变量回归vs多元回归

根据自变量的个数,线性回归分为两类:

单变量回归:只有一个自变量的情况

多元回归:多于一个自变量的情况

2.3 线性回归的数学表达

线性回归试图学得一个线性模型,以尽可能准确地预测实值输出标记。

给定数据集 ,其中 ,是第个样本的个特征,是对应的真实标记。

线性模型的一般形式为:

写成向量形式:

其中是权重向量,是偏置项。学得之后,模型就确定了。

2.4 更深刻地理解线性回归

我们来看几个具体的例子:

例1:第N天发的硬币数量

第N天发的硬币数量=1*天数

这是一个一元线性模型,只有一个特征(天数)。

例2:期末成绩

期末成绩=0.7*考试成绩+0.3*平时成绩

这是一个多元线性模型,有两个特征(考试成绩、平时成绩)。

例3:房子价格

房价= 0.02*中心区域距离+0.04*一氧化氮浓度+(-0.12)*自住房平均房价+0.254*城镇犯罪率

这也是一个多元线性模型,有四个特征。

从这三个例子可以看到,特征值与目标值之间建立了一个线性关系,这个关系就是线性模型。每个特征前面的系数代表了该特征对目标值的 "贡献程度"。

3. 线性回归的特征与目标关系分析

线性回归当中主要有两种模型关系:线性关系非线性关系

3.1 线性关系

3.1.1单变量线性关系

当只有一个特征时,特征与目标值的关系呈直线关系 。在二维坐标系中,就是一条斜率为、截距为的直线。

3.1.2多变量线性关系

当有两个特征时,特征与目标值呈现平面关系 。在三维坐标系中,就是一个平面。推广到个特征,就是维空间中的一个超平面。

3.2 非线性关系

如果特征与目标值之间不是简单的线性关系,那么回归方程可以理解为更复杂的形式。但需要注意的是,线性回归中的 "线性" 是指参数是线性的 ,而不是指特征必须是线性的。

例如,虽然对来说是非线性的,但对参数来说仍然是线性的,因此仍然属于线性回归的范畴。

4. 线性回归的应用场景

线性回归是机器学习中最基础、应用最广泛的算法之一,常见的应用场景包括:

房价预测:根据房屋面积、位置、房龄等特征预测房价

销售额度预测:根据广告投入、促销活动、季节等因素预测销售额

贷款额度预测:根据收入、信用评分、负债情况等预测可贷款额度

除此之外,线性回归在金融、经济、医学、社会科学等领域都有广泛应用。它不仅可以直接用于预测,还可以用来分析各个特征对目标值的影响程度。

5. 如何建模:寻找那条最合适的线

回到发硬币的例子,当数据点不在一条直线上时,我们需要找到一条直线,让它尽可能地 "拟合" 这些数据点。

所谓 "建模",就是要确定线性模型中的参数(斜率)和(截距)。

那么问题来了:什么样的线才算 "最合适"?

直观上,我们希望所有数据点到这条直线的 "距离" 尽可能小。这就引出了损失函数的概念。

6. 损失函数:如何量化 "拟合得好不好"

6.1 残差与残差平方和

对于每个样本点,模型的预测值为。预测值与真实值之间的差异称为残差

为了衡量整体拟合效果,我们把所有残差的平方加起来,得到残差平方和(Residual Sum of Squares, RSS)

为什么要用平方而不是绝对值?因为平方可以放大较大的误差,让模型更关注那些偏离较大的点;同时平方函数是光滑可导的,便于后续的优化计算。

6.2 均方误差与损失函数

我们的目标是让均方误差最小化 ,即找到最优的,使得:

代入线性模型,可写为:

注意:在求最优参数的表达式中,前面的常数系数(如)不影响最优解的位置,因此可以省略。

而损失函数(Loss Function)本身的定义(带均方误差的系数)为:

其中对应斜率对应截距。损失函数也叫代价函数(Cost Function),它的值越小,说明模型拟合得越好。

补充:在很多机器学习教材(如吴恩达的课程)中,损失函数还会在前面乘以,即,这个是为了在求导时与平方项的抵消,使梯度表达式更简洁,不影响最优解。

在线性回归中,基于均方误差最小化来进行模型求解的方法称为最小二乘法(Least Squares Method)

7. 线性回归的求解方法

线性回归的损失函数是一个二元凸函数 (对于单变量线性回归,参数是,这是一个无条件极值问题。凸函数的极值点就是最小值点,而极值点处梯度为0,因此可以转化为求解多元线性方程组。

线性回归中经常使用的两种求解方法:

7.1 最小二乘法(正规方程法)

最小二乘法通过令损失函数对每个参数的偏导数为0,解线性方程组直接得到最优参数的解析解。

对于多元线性回归,写成矩阵形式后,最优参数为:

其中是样本特征矩阵,是真实标记向量。

最小二乘法的局限性 :在真实任务中,特征矩阵往往不是满秩矩阵。当数据集有大量特征,特征数目甚至超过样本数时,的列数多于行数,显然不满秩,根本无法求逆矩阵。这时候最小二乘法就失效了。

7.2 梯度下降法

一个更好的处理方式是使用梯度下降算法来求取最优值。梯度下降法是一种迭代优化算法,不要求矩阵可逆,适用于各种复杂的优化问题。

注意:在真正的开发过程中,梯度下降法使用最多,在深度学习中更加明显。

8. 梯度下降算法详解

8.1 什么是梯度

在理解梯度下降之前,我们首先要明白什么是梯度

梯度:

在单变量函数中 ,梯度其实就是函数的微分,代表着函数在某个给定点的切线的斜率

在多变量函数中 ,梯度是一个向量 ,向量有方向,梯度的方向就指出了函数在给定点的上升最快的方向

用微积分的语言来说,对多元函数的参数求偏导数,把求得的各个参数的偏导数以向量的形式写出来,就是梯度。

例如,对于二元函数,其梯度为:

8.2 梯度下降的直观理解

想象自己站在一座山上的某个点,你想要下山到谷底。环顾四周,沿着最陡峭的方向挪一小步,到达一个新的点;继续环顾四周,找到当前最陡峭的方向,再挪一步...... 如此反复,最终就能到达谷底。

梯度下降就是这样一个过程:从某个初始点出发,每次沿着当前点梯度的反方向(因为梯度方向是上升最快的方向,反方向就是下降最快的方向)移动一小步,逐步逼近函数的最小值点。

用一句话解释:梯度下降法就是快速找到函数最低点的一个方法。就像山上有一个球,经过几次滚动后,就会来到谷底附近。

9. 梯度下降法三要素:方向、距离、终止条件

梯度下降法的完整过程需要确定三个关键要素:

9.1 方向:确定往哪个方向滚

梯度的方向是函数上升最快的方向,那么梯度的反方向就是函数下降最快的方向。因此,每次迭代时,我们沿着当前点梯度的反方向前进。

这就是为什么梯度下降公式中梯度前面要乘以一个负号------我们要朝着与梯度相反的方向前进,也就是朝着下降最快的方向走。

9.2 距离:确定滚多远

每次移动多远呢?这由学习率(Learning Rate)来控制,通常用表示。学习率决定了每一步走的距离大小。

9.3 终止条件:确定滚到哪里才算结束

我们不可能无限迭代下去,需要设定一个终止条件。常见的终止条件包括:

梯度值下降到某个阈值以下(如梯度的模长

损失函数的变化量小于某个阈值

达到预设的最大迭代次数

10. 学习率对梯度下降的影响

学习率(也记作)是梯度下降中最重要的超参数之一,它的选择直接影响算法能否收敛以及收敛的速度。

为了直观地对比不同学习率的效果,我们统一使用以下设定:

目标函数:(最小值在处)

梯度:

起始位置:

迭代公式:

10.1 学习率太小:

当学习率过小时,每一步移动的距离非常小。我们来看看前几次迭代的结果:

可以看到,每次迭代只向最低点移动了一点点(从10到9.6只移动了0.4),导致收敛速度极慢,需要迭代很多次才能到达最低点。虽然最终能收敛,但效率很低。

10.2 学习率合适:

当学习率选择合适时,算法能够以较快的速度稳定收敛到最小值点。

开始,迭代过程如下:

可以看到,的值在快速减小,稳步向最低点逼近,没有出现震荡或发散。

对应的梯度值变化为:

|--------------------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|--------------------------------------|
| 迭代次数 | | | | | | | | | | | |
| | 20 | 12 | 7.2 | 4.32 | 2.59 | 1.56 | 0.93 | 0.56 | 0.34 | 0.2 | 0.12 |

梯度值每次乘以,呈指数级下降,收敛稳定且高效。

10.3 学习率太大:

当学习率过大时,每一步移动的距离太大,会跨过最低点 ,在最低点两侧来回震荡,而且震荡幅度越来越大,最终导致发散

我们来看看前几次迭代的结果:

可以清楚地看到:

(跨过了最低点,到了左边)

(又跨过最低点,到了右边,而且比上次更远)

(再次跨过,更远了)

(越来越远!)

数值的绝对值越来越大,完全无法收敛到最低点。

重要提醒 :完成梯度下降,必须选择合适的学习率!不能太大也不能太小------太小会导致迟迟走不到最低点,太大会导致错过最低点甚至发散。

在实际应用中,通常会尝试多个学习率值(如 0.001、0.01、0.1、0.2、0.5等),观察损失函数的下降曲线来选择最合适的学习率。

11. 梯度下降的终止条件

我们以学习率、函数、起始位置为例,来看看终止条件是如何设定的。

11.1 每次迭代后的梯度值

首先计算出学习率为时,每次迭代后的梯度值

|--------------------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|-----------------------------------|--------------------------------------|
| 迭代次数 | | | | | | | | | | | |
| | 20 | 12 | 7.2 | 4.32 | 2.59 | 1.56 | 0.93 | 0.56 | 0.34 | 0.2 | 0.12 |

可以看到,随着迭代次数的增加,梯度值在不断减小:从初始的20,到第10次迭代时已经降到了 0.12。当梯度值接近0时,说明已经非常接近函数的最小值点(此时)。

11.2 用梯度值作为终止条件

既然梯度值会随着迭代不断减小,我们就可以设定一个阈值,当梯度值小于等于这个阈值时,就认为已经足够接近最小值点,可以停止迭代了。

案例 1 :希望最后的梯度值小于等于,即

由于每次迭代后梯度值乘以,即:

,解得。因此,只要迭代15次,梯度值就会降到0.01以下,任务完成!

案例 2 :如果要求更严格,希望梯度值降到更小的阈值(例如),那么就需要迭代19次才能满足条件。

由此可见,终止条件的阈值设定直接影响迭代次数:阈值越小,需要迭代的次数越多,结果也越精确,但计算成本也越高。

11.3 理论保证

只要学习率选择得合适,梯度就可以下降到任意小。这是有理论支撑的(梯度收敛定理):对于凸函数,当学习率满足一定条件时,梯度下降算法保证收敛到全局最小值,且梯度值可以任意接近0。

有了这个理论保证,我们就可以放心地选择梯度值作为终止条件

11.4 常见的终止条件设定方式

在实际应用中,常见的终止条件有以下几种:

梯度的模长 (如):最常用的方式,当梯度足够小时停止

损失函数变化量:连续几次迭代损失函数的变化量小于某个阈值,说明损失已经收敛

最大迭代次数:设定一个迭代上限(如1000次),作为安全兜底,防止因学习率不合适导致无限循环。

通常会将多种终止条件结合使用,例如 "梯度值或迭代次数达到1000次时停止",既保证精度又防止死循环。

12. 多元函数的梯度下降

前面我们主要以单变量函数为例讲解了梯度下降的原理,现在通过一个具体的二元函数例子,来看看多元函数的梯度下降是如何一步步迭代的。

12.1 问题设定

设函数:

起始位置为,在图中用红点表示,学习率为

这个函数是一个典型的凸函数,其最小值在原点处,最小值为。我们的目标就是通过梯度下降法从起始点出发,逐步逼近这个最小值点。

12.2 计算梯度

首先求函数的梯度。梯度是一个向量,由各个偏导数组成:

在起始位置处,梯度为:

12.3 第一次迭代

梯度下降的迭代公式为:

将起始位置和梯度代入,进行第一次迭代:

所以第一次迭代后,位置从移动到了,离原点更近了一步。

12.4 后续迭代与终止条件

根据迭代公式,我们可以计算出每次迭代后的位置。每一次迭代都重复以下步骤:

计算当前位置的梯度

沿梯度反方向移动一步:

将终止条件设为梯度的模长小于等于 0.01 ,即

经过不断迭代,到第 16 次时:

此时梯度已经非常小,说明已经非常接近最小值点,任务完成!

注意 :在更新参数时,必须同时更新 所有参数(即使用上一次迭代的梯度值来计算所有新参数),不能更新一个参数后就用新值去计算另一个参数的梯度。例如第一次迭代中,都是基于处的梯度计算的,而不是先算再用去算

13. 梯度下降法总结

13.1 梯度下降公式

梯度下降的核心公式为:

其中:

是第个参数

是学习率(步长)

是损失函数对第个参数的偏导数

13.2 学习率的含义

在梯度下降算法中被称作学习率 或者步长 ,意味着我们可以通过来控制每一步走的距离。

太小 :可能导致迟迟走不到最低点,收敛缓慢

太大:会导致错过最低点,甚至无法收敛

选择合适的学习率是梯度下降成功的关键

13.3 为什么梯度要乘以一个负号?

梯度前加一个负号,就意味着朝着梯度相反的方向前进

梯度的方向实际就是函数在此点上升最快的方向 ,而我们需要朝着下降最快的方向走,自然就是负的梯度的方向,所以此处需要加上负号。

13.4 梯度下降优化过程演示

整个优化过程可以概括为:

1.随机初始化参数

2.计算当前点的梯度

3.沿梯度反方向更新参数:

4.检查是否满足终止条件,若不满足则返回步骤2;

5.满足终止条件,得到最优参数。

通过不断迭代,参数会逐步逼近损失函数的最小值点,模型的预测效果也会越来越好。

写在最后

线性回归是机器学习的入门算法,也是理解更复杂算法的基础。它的核心思想非常朴素:用一个线性函数来拟合数据,通过最小化损失函数来找到最优参数

而梯度下降法则是机器学习中最核心的优化算法之一,不仅用于线性回归,更是深度学习的基石。理解梯度下降的原理(方向、步长、终止条件),对于学习后续的算法至关重要。

希望这篇总结能帮助大家更好地理解线性回归和梯度下降。如有错误或遗漏,欢迎指正交流!

相关推荐
HugoStudio_SWAN1 小时前
洛谷 P1420 / P1179 / B4262 最长连号、数字统计与词频统计——统计的三种面孔
c++·学习·程序人生·算法
天远Date Lab1 小时前
零信任架构实战:基于天远车信盟出险构建自动化汽车消费贷合规网关
人工智能·机器学习·计算机视觉·ocr
是隼人1 小时前
buuctf-pwn [HarekazeCTF2019]baby_rop2(64位ret2libc)题解(学习过程持续更新)
c语言·学习·安全·pwn入门·ctf入门
阳明山水2 小时前
销量预测2025—2026:从模型竞赛到系统融合的范式转型
人工智能·深度学习·算法·机器学习·架构
gzkuijing3 小时前
MDI1PRD23B7‑EQ|Novanta IMS一体化RS-485/422步进电机参数、选型要点与典型应用
人工智能·机器学习·编辑器
aprilaaaaa3 小时前
(TC397 学习笔记)二、ICU和TIM的时钟计算
笔记·学习·fpga开发
147API3 小时前
蒸馏项目什么时候该停,怎样切换到RAG、微调或模型路由
人工智能·深度学习·机器学习·数据挖掘
xian_wwq4 小时前
【学习笔记】-深度认知系列-第8讲主流AI模型横向评测——GPT、Claude、Gemini、盘古、文心、通义
人工智能·笔记·学习
大牧师4 小时前
TypeORM 学习教程
数据库·sql·学习·node.js·orm·nest.js·typeorm