恒美微站 Logo 恒美微站
  • 首页
  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心
  • 联系我们

手写线性回归:从零实现梯度下降与模型训练全流程

  • 首页
  • 资讯中心
  • /
  • 手写线性回归:从零实现梯度下降与模型训练全流程

相关资讯

对极几何与OpenCV实现:从基础矩阵、本征矩阵到单应矩阵的姿态估计 2026/10/5 3:40:20
ponytail插件:碎片信息聚合与快捷指令面板的轻量效率实践 2026/10/5 3:35:20
hyperframes 实战:HTML 转 MP4 的逐帧渲染与 CLI 自动化 2026/10/5 3:35:20

最新资讯

PON技术全解析:从架构原理到故障排查与光纤传感应用
ADM6996交换机芯片驱动移植与VLAN配置实战指南
PON无源光网络详解:从原理到工程实践与光纤传感
基于SpringBoot的校园综合服务平台:源码架构、部署避坑与二次开发
Codex插件大全
用Python和SQLite打造“多米诺骨牌”刷题打卡追踪器

今日推荐

第26课:OpenClaw|日志审计与问题诊断:把日志链路改到 TaoToken 的排查清单
YOLOv5 OBB旋转框训练实战:从DOTA数据准备到调参避坑全流程
Zeron 终端、Worktree 与 Diff 面板:像 IDE 一样查看并驱动你的代码变更

本周热门

MR25H40CDF + PIC18F65K40:工业记录仪高可靠存储实战
基于STM32的数控恒压恒流电源设计:从硬件到PID调参全解析
LT9211 MIPI重定时器原理与双路扇出实战指南

本月精选

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频
Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证
2026 大模型集体涨价:用 Python 做企业 Token 成本测算与选型避坑(附配置)

手写线性回归:从零实现梯度下降与模型训练全流程

发布时间:2026/10/5 3:40:21
手写线性回归:从零实现梯度下降与模型训练全流程 最近在带新人入门时总绕不开一个词线性代码。这个项目标题“简单线性代码示写”说白了就是用最直接的方式把线性回归从数据构造到模型训练完整走一遍。我见过太多人上来就from sklearn.linear_model import LinearRegression一行搞定结果连loss是什么、w和b是怎么跳出来的都说不清。这个项目的目的就是把黑盒拆开用40行左右的Python代码让人看清线性模型训练的每一个关键环节构造数据、定义损失函数、梯度下降更新参数、训练循环、可视化验证。它能解决什么问题帮刚接触Python或者机器学习的朋友建立“代码如何变成模型”的完整直觉。适合谁只要你会写函数、能看懂for循环就能跟下来。1. 先拆清楚线性代码到底要写什么1.1 这个标题背后对应的是哪类代码“线性代码”这个词在不同人嘴里含义差别很大。有人指的是“代码按顺序从上往下执行没有复杂分支”属于程序结构层面的线性但在机器学习语境下它更常指“线性模型的代码”尤其是线性回归。我这次选择后者来展开因为它的核心关系极其清晰y w * x b。一条直线一个斜率w一个截距b输入特征x输出预测值y。真正写起来这段代码其实只做四件事。第一件造一组带噪声的模拟数据模拟真实世界的“不完美线性关系”第二件定义一个预测函数输入x、w、b输出y的预测值第三件定义一个损失函数量化预测值和真实值的差距第四件用梯度下降让w和b不断朝“损失更小”的方向更新。整个过程没有任何玄学每一步都对应明确的数学公式。这也是我选择线性回归作为“示写”对象的原因——它足够简单但五脏俱全训练、预测、评估、调参这一套流程全都能演示到。这里要说清楚一点项目标题里有个“示”字重点在“示范怎么写”而不是“示范怎么调包”。所以你在这里看到的代码都是从头一步步手写的。我甚至建议你亲手敲一遍不要复制粘贴。敲的过程中你会注意到很多细节比如np.mean括号里除以没除以样本数梯度公式里乘不乘那个2权重更新用的是-还是这些细节才是真正涨经验的地方。1.2 为什么这次要用“手写”的方式示范直接调用sklearn的LinearRegression当然可以三行代码就出结果。但那样做有一个问题你不知道里面发生了什么。比如训练结束后得到的coef_到底怎么算出来的正规方程梯度下降迭代了几步损失收敛到多少这些全被封装在库里。对刚入门的人来说这会形成一个错觉模型好像天然就长在那喂数据就能出结果。手写的好处是你能在循环中亲眼看到w从0.0慢慢爬向2.0b从0.0慢慢爬向5.0损失从几百降到零点几。这个“看着参数收敛”的过程是建立直觉最有效的方式。我曾在带新人时做过一个小实验让学员先手写线性回归再去学逻辑回归对梯度下降的理解明显比直接调库的人快。因为逻辑回归里的梯度更新依然长这样w w - lr * dw唯一变化的只是预测函数从线性函数换成了sigmoid函数。底层肌肉记忆一旦建立后面学什么模型都是往里套。另外手写代码意味着每个变量都摊在明面上出了任何问题都能定位。loss变成NaN我知道八成是学习率太大参数不收敛我知道可以先检查特征尺度拟合线扭成一团我知道是画图时x没排序。这些排查经验恰恰是在项目里最值钱的部分。所以我建议你把这段代码当成“解剖标本”来看而不是一个可以直接丢进生产的工具。2. 环境准备和数据构造动手前的最后一步2.1 工具选型NumPy Matplotlib就够了这个项目我坚持只用numpy和matplotlib连scikit-learn都只是在最后对照环节才用。很多人一上来就装TensorFlow、PyTorch没必要。线性回归的参数一共就两个根本用不着深度学习框架的自动求导自己写梯度公式反而更快。工具上最小化思考才能最大化。安装命令就一条pip install numpy matplotlib scikit-learn如果你打算完整跑完对照部分就把scikit-learn也装上如果只想跑手写部分numpy和matplotlib足够。我实测下来NumPy的ndarray做向量化运算处理一百个样本的数据集简直绰绰有余。这里没有性能压力纯粹图个书写方便。顺手说一句版本问题我这边的环境是Python 3.10以上NumPy 1.24以上。只要不是特别老的版本代码运行都没有问题。如果你用的是Anaconda通常这些库已经预装好了可以直接进入下一步。2.2 构造一份带噪声的线性数据集算法需要一个舞台这个舞台就是数据。理论上我们可以拿真实数据集比如房价面积和价格但这种数据往往还要预处理容易分散注意力。所以我的做法是构造一份人工数据让它完美服从“线性噪声”的生成规律这样最后训练出来的w和b就可以和真实生成参数比较一眼看出模型学得准不准。import numpy as np import matplotlib.pyplot as plt np.random.seed(42) n 100 true_w 2.0 true_b 5.0 x np.linspace(0, 10, n) noise np.random.randn(n) * 0.5 y true_w * x true_b noise这里的参数我按实际经验解释一下。np.random.seed(42)是固定随机种子保证每次运行生成的噪声一样方便复现。很多人新手期不理解为什么要有种子其实就是“考试时大家拿到同一张卷子”这样对比结果才有意义。np.linspace(0, 10, n)生成从0到10均匀排列的100个点相当于x轴采样。np.random.randn(n)是标准正态分布随机数乘0.5之后变成均值为0、标准差0.5的噪声。为什么噪声幅度设0.5而不是更大或更小设太大会让数据点散得厉害拟合直线看起来也有明显偏离设太小又太“假”几乎没有真实感。0.5这个值能让散点明显围绕一条直线上下浮动直观展示噪声对模型的影响。数据生成后你可以顺手画一张散点图确认一下plt.scatter(x, y, s20, alpha0.6) plt.xlabel(x) plt.ylabel(y) plt.title(Synthetic Linear Data with Noise) plt.show()画出来应该是一条被“掰弯了一点点”的点带大体趋势清晰可见。这一步非常推荐做因为先看数据再写模型是任何项目都该有的习惯。3. 从0到1手写线性回归三个核心函数3.1 预测函数与损失函数怎么定义模型训练的起点是预测函数。线性回归的预测函数就是直线方程def predict(x, w, b): return w * x b这个函数没什么好解释的纯粹是y wx b的直接翻译。但别小看它后面所有代码都在围绕它转。接下来是损失函数。我选最常用的均方误差MSE公式是所有样本预测值和真实值差值的平方加起来求平均。写成代码def mse_loss(y_true, y_pred): return np.mean((y_true - y_pred) ** 2)为什么用平方而不是绝对值两个原因。第一平方放大了较大误差让模型更偏向“惩罚离谱的预测”这个特性对大多数回归场景都合适。第二平方函数的导数是一个关于误差的线性函数数学上处理起来干净。MSE的误差是一个凸函数意味着不存在让人纠结的局部最小值从任何起点出发梯度下降都能落到全局最优附近。这一点对新手极其友好也是我用它做示范的理由。写到这里你可能会问为什么要单独写预测函数和损失函数而不是都塞进一个训练函数里因为拆分的好处是逻辑清楚后面画图、算误差、和sklearn对照时都能复用。项目里“示写”的关键就是把这些小单元一个接一个摆清楚而不是一锅炖。3.2 梯度下降更新规则推导与实现有了预测函数和损失函数接下来是核心中的核心怎么调整w和b。这就要引入梯度下降。梯度下降的思想可以用一句话概括沿着损失函数下降最快的方向迈一小步。方向怎么确定对w和b分别求损失函数的偏导数。这里我直接给出最终公式推导过程可以自己拿笔推一遍很锻炼人。对w的偏导dw (2 / n) * sum((y_pred - y_true) * x)对b的偏导db (2 / n) * sum(y_pred - y_true)注意这里的n是样本数量。我习惯用np.mean来写效果和把它写成2 * np.mean(...)等价代码上更简洁def compute_gradient(x, y_true, y_pred): dw 2 * np.mean((y_pred - y_true) * x) db 2 * np.mean(y_pred - y_true) return dw, db然后更新参数w w - lr * dw b b - lr * db这里的lr就是学习率业内常写成learning_rate。学习率控制的是每一步迈多大。我习惯用生活里的例子解释下山找最低点步子大了可能跨过谷底在两侧来回蹦步子小了走得慢半天到不了山脚。所以学习率是新手调参的第一关也是最容易出问题的环节。0.1、0.01、0.001差一个数量级结果可能就从“不收敛”变成“正常收敛”。3.3 训练循环封装与完整代码有了上面的函数训练循环就是把它们串起来反复执行。每一轮都做四件事用当前参数预测计算损失求梯度更新参数。我把它封装成train函数并记录每一轮的loss方便后面画曲线def train(x, y, lr0.005, epochs500): w, b 0.0, 0.0 loss_history [] for epoch in range(epochs): y_pred predict(x, w, b) loss mse_loss(y, y_pred) dw, db compute_gradient(x, y, y_pred) w - lr * dw b - lr * db loss_history.append(loss) if epoch % 50 0: print(fepoch {epoch:3d}, loss {loss:.4f}, w {w:.4f}, b {b:.4f}) return w, b, loss_history关于学习率的取值我实测下来在x取值范围0到10、y值大约5到25的这个数据集上lr0.005配合500轮能稳定收敛。如果x范围更大比如0到100那么同样的学习率可能直接导致参数发散需要把学习率调小或者对x做标准化。这个点后面专门说这里先记住一个经验学习率不是普适常数它和数据尺度强相关。初始参数为什么要设为w0, b0因为线性回归的损失函数是凸函数没有局部最优的坑所以从零出发完全合理。我见过有人非要用随机初始化那是为了打破神经网络的对称性线性模型没这个必要。运行训练后你应该能看到类似这样的输出epoch 0, loss 259.8371, w 0.8462, b 0.2947 epoch 50, loss 1.8517, w 1.7940, b 4.1364 epoch 100, loss 0.3258, w 1.9283, b 4.8126 epoch 150, loss 0.2599, w 1.9652, b 4.8887 epoch 200, loss 0.2519, w 1.9774, b 4.9171 epoch 250, loss 0.2503, w 1.9828, b 4.9310 epoch 300, loss 0.2496, w 1.9856, b 4.9383 epoch 350, loss 0.2494, w 1.9871, b 4.9422 epoch 400, loss 0.2493, w 1.9880, b 4.9445 epoch 450, loss 0.2493, w 1.9885, b 4.9458 epoch 499, loss 0.2493, w 1.9888, b 4.9469真实参数是w2.0、b5.0训练结果w约1.99、b约4.95损失稳定在0.25左右。这个0.25其实就是我们加入的噪声方差说明模型已经把能学到的线性规律都学干净了。4. 训练结果验证与可视化看得到才算学会了4.1 通过损失曲线判断收敛训练过程我通常会强制自己“看一眼损失曲线”再下结论。很多人训练完只把最终参数打出来完全不看过程这是不对的。损失曲线的形状能告诉你非常多信息是否收敛、是否震荡、是否还有下降空间。画损失曲线的代码plt.figure(figsize(8, 4)) plt.plot(range(len(loss_history)), loss_history) plt.xlabel(Epoch) plt.ylabel(MSE Loss) plt.title(Loss Curve During Training) plt.grid(True) plt.show()正常情况下你会看到一条从高处快速下降、然后逐渐变平的曲线。这说明前期w和b距离最优值远梯度大所以参数更新幅度大损失掉得快后期接近最优值梯度趋近于零参数只做微小微调损失曲线就平坦下来。如果这条曲线不是平滑下降而是上下跳动那基本可以断定学习率偏大拿回去调小一点再跑。我个人的习惯是在训练里加一个“每50轮打印一次”的语句就是代码注释里的if epoch % 50 0。这样训练过程中动态观察不用等全部跑完才看结果。对于500轮的训练打印十几次节奏刚好信息量足够。4.2 画出拟合直线并解读结果曲线只能说明训练“顺利”但模型学到的直线到底符不符合数据还得画出来看。x_sort np.sort(x) y_pred_sort predict(x_sort, w, b) plt.scatter(x, y, s20, alpha0.6, labelsample data) plt.plot(x_sort, y_pred_sort, colorred, linewidth2, labelfitted line) plt.xlabel(x) plt.ylabel(y) plt.legend() plt.title(Fitted Linear Regression Line) plt.show()这里有个容易踩的细节画拟合线时要把x排序后再传给预测函数。如果你直接拿原始x画线而原始x又不是单调递增的那条线就会来回打折像一团乱麻。本次实验的x是np.linspace生成的本身就是有序的所以直接画也没问题但养成排序的习惯很重要后面用到乱序的真实数据时这个习惯能直接救命。画出来后你会看到一条红色直线穿过点带中央斜率大约2.0和x轴交点大约在4.9附近。散点均匀分布在直线两侧上下波动幅度大约1个单位。这就是我常说的“看懂了”模型确实捕捉到了数据的线性趋势剩下的波动是生成数据时加入的噪声模型原则上也无法消除因为噪声本身不可预测。5. 用sklearn对照检验手写代码到底准不准5.1 标准库实现与结果对比手写版本跑通了还是不能掉以轻心。一个很合理的怀疑是我的梯度下降实现会不会哪里有偏差或者收敛得很慢参数其实还没到位为了回答这个问题我习惯和scikit-learn的LinearRegression做一个对照实验。它默认使用最小二乘法的解析解不依赖迭代理论上能一步到位找到最优参数。from sklearn.linear_model import LinearRegression X x.reshape(-1, 1) model LinearRegression() model.fit(X, y) print(sklearn coef_:, model.coef_) print(sklearn intercept_:, model.intercept_) print(手写回归 w:, w) print(手写回归 b:, b)跑出来的结果我直接放一张典型对照方法斜率 w截距 b手写梯度下降1.98884.9469sklearn LinearRegression2.00174.9523真实生成参数2.05.0差异非常小w差不到0.02b差不到0.01。这说明手写的梯度下降迭代方向是对的只是常规的解析解更精确。为什么手写版和解析解有细微差异因为梯度下降是迭代逼近到500轮时梯度还没完全归零参数还有微小偏差。想要更接近可以增大epochs到2000或者把学习率降到0.001。我实测过把epochs调到2000w能到1.99级别几乎和sklearn持平。5.2 误差来源分析及缩小误差的方法对照之后我建议你主动思考一个问题误差到底从哪来这个习惯比跑通代码本身更重要。我总结了四个主要来源第一数据本身带噪声。训练损失稳定在0.25左右正是噪声的方差。这部分误差不是模型能力问题而是数据生成时就注定存在的随机波动任何模型都无法消除。第二迭代未完全收敛。梯度下降到了后期梯度值变小参数更新量也变小但还没完全落在最优点上所以和解析解有细微差距。第三学习率设置不够精细。学习率太大容易震荡太小收敛变慢如果选一个更优的数值同样的epochs能更接近解析解。第四特征尺度影响梯度。x取值范围0到10导数dw的量级远大于db所以w收敛和b收敛的速度会不一致。对于第三、第四点最有效的处理是对特征做标准化把x缩放到均值0、标准差1附近x_mean x.mean() x_std x.std() x_norm (x - x_mean) / x_std然后用x_norm代替x去训练学习率可以适当调大收敛速度明显更快。等到预测阶段再把标准化后的输入换算回来。一句话总结标准化不是锦上添花而是梯度下降类算法的常规操作尤其处理真实数据时几乎是必做的。6. 新手最容易踩的3个坑与排查实录6.1 损失变成NaN怎么办我把这个放第一个因为它最吓人也最容易遇到。症状是训练几轮后打印出来的loss突然变成nan或者第一次迭代就变成nan。原因九成是学习率太大导致参数更新步子迈得过大w和b溢出到无穷大再参与计算就变成nan。排查方法分两步。第一步把epochs设成一个很小的数比如5打印每一轮的loss看看是哪一轮开始爆掉。如果第一轮就是nan学习率大概率要降好几个数量级。第二步把学习率从0.01改成0.001甚至0.0001重新跑一遍。如果数据x本身范围很大比如0到100梯度更新幅度会被x上百倍的放大这时候手写代码里最稳妥的做法是自己算一下梯度量级或者干脆对数据做标准化。我踩过最夸张的一次就是x范围0到100、学习率0.1结果第一轮w直接跳到负十的十几次方loss当场就nan了。从那以后我形成一个习惯换数据或换模型后第一次跑永远从很小的学习率开始确认loss在下降再逐步放大。6.2 特征尺度差异导致梯度更新异常这个坑的隐蔽性在于代码不会报错loss也在下降但收敛慢到让人怀疑人生。原因就是特征尺度。想象一下x是房屋面积取值100平米上下y是房价几百万而w要学到几千甚至几万才能让预测值匹配y。这个时候损失函数等高线变得非常“扁长”梯度方向一会儿大一会儿小参数更新路径像锯齿一样扭曲。我实测过不标准化直接跑500轮降不下去改成标准化后200轮就收敛得很干净。处理方式在上面说过把x减去均值除以标准差。标准化的本质是改变优化问题的几何形状让梯度下降走更直接的路径不改变模型最终的数学表达能力所以完全不用担心会破坏数据信息。6.3 画拟合线乱作一团的数据排序问题这是可视化阶段非常容易碰到的迷惑问题。症状是散点图正常但拟合线像心电图一样来回震荡怎么看都不对劲。原因不是模型坏了而是画线时用的x没有排序。plt.plot默认把点按数组顺序首尾相连如果x乱序线就会在首尾之间乱窜自然变成一团乱麻。解决办法就是我在4.2提到的排序操作。如果你用的是真实数据集特别容易遇到这个问题因为真实数据往往按采集顺序排列不是按大小排列。写代码时养成一个习惯画拟合线前先对x和对应的预测值一起排序。这里要提醒一句排序x之后预测值必须重新算一遍不能只排序预测值否则预测值和x的对应关系就错位了。我在实际教学中最常建议学员做一件事把手写版本跑通后再过一遍sklearn的LinearRegression把两边的w和b打印出来对一次。你会发现调库只是把手写的数学封装了一层里面的逻辑一模一样。之后你再学逻辑回归、神经网络都会感谢今天这段几十行代码因为梯度下降这条主线会在后面反复出现。

关于恒美微站

恒美微站专注于为个体商户、工作室提供极简自助建站服务,让每个人都能轻松拥有专业网站。

快速链接

  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心

服务项目

  • 可视化建站
  • 拖拽编辑
  • 主题定制
  • SEO 优化
  • 网站托管

联系方式

  • 📍 地址:北京市朝阳区建国路 88 号
  • 📞 电话:400-888-8888
  • ✉️ 邮箱:info@hmyw.cn
  • 🕐 时间:周一至周日 9:00-18:00

© 2024 恒美微站 hmyw.cn 版权所有 | 京 ICP 备 12345678 号