ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

最简单的线性回归代码:手写梯度下降与sklearn实战

最简单的线性回归代码:手写梯度下降与sklearn实战 先纠正一个小细节liner大概率是linear线性的笔误这个词在搜索框里出现的频率还挺高很多人写代码时也会顺手打错。这篇笔记是最简单的回归代码系列的第4篇前几篇拆过理论、讲过公式推导这篇集中火力只聊代码本身怎么用最短的代码把线性回归跑起来手写和调库各来一遍参数怎么设、数据怎么预处理、跑出结果怎么判断好坏。适合谁看刚上手机器学习、被各种教材里最小二乘梯度下降绕晕的新手已经会用sklearn但想搞清楚底层逻辑的进阶者以及想快速把回归模型用在自己数据上的朋友。这篇尽量照顾到不同基础能抄代码的地方直接抄想深挖原理的地方也有解释。1. 这篇笔记在聊什么一个纯粹的跑通问题1.1 先搞明白最简单到底指什么很多人第一次写回归代码最大的障碍不是公式而是代码不知道从哪一行开始写。线性回归作为入门算法它的代码量可以压缩到非常小手动实现核心逻辑也就二十几行调库的话三行就够。所谓最简单是指它的数学假设最简单——因变量和自变量之间是线性关系一条直线或一个超平面去拟合也是指代码结构最简单——没有复杂的网络层、没有繁琐的预处理流程是把Python NumPy 一份数据组合起来就能跑的最小闭环。但要提醒一句代码短不代表可以闭眼抄。我在实际写代码过程中发现线性回归的代码陷阱往往不在回归本身而在数据。数据有没有空值、量纲差多少、有没有明显的异常点这些直接决定你跑出来的系数是可解释的还是看起来对但没卵用的。所以这篇笔记表面上写代码实际上把数据准备和结果验证也一并讲透。1.2 为什么是第4篇学习笔记这个系列前几篇大概覆盖了什么是回归问题、误差怎么衡量、正规方程和梯度下降的推导。到了第4篇默认读者已经知道损失函数是均方误差目标是找到让损失最小的w和b。如果你还没看前几篇也没关系这篇涉及到的公式都会用代码重新演示一遍——代码本身就是最好的公式解释器。我始终觉得学机器学习有个笨但特别有效的路径先照着敲敲完改参数改完再看结果变化最后再回头推公式。这比上来啃一堆矩阵求导舒服得多也记得牢得多。2. 写代码之前的三件事环境、数据、损失函数2.1 环境依赖其实就三个跑线性回归不需要深度学习框架有Python就够。我建议的环境是Python 3.8及以上低版本不是不行但3.8以下有些新语法和库版本不好适配NumPy所有矩阵运算的底子Matplotlib画散点图和回归线一定要学因为判断线性拟合好不好肉眼比指标快scikit-learn调库版回归要用但手写代码阶段用不到它安装命令不写了网上一搜一大把。重点提醒一下scikit-learn的版本问题0.24之后和1.x版本的API有细微变化如果你下载的教程代码报AttributeError先检查是不是版本不对。我踩过这个坑一度以为是代码写错了最后发现是版本混用。2.2 造一份能用来练手的数据没有真实数据的时候自己造数据是最快的验证方式。我一般用这种形式import numpy as np import matplotlib.pyplot as plt np.random.seed(42) X np.random.rand(100, 1) * 10 # 100个样本特征取值范围0~10 true_w, true_b 3.0, 5.0 # 真实系数y 3x 5 y true_w * X.squeeze() true_b np.random.randn(100) * 2 # 加一点噪声这段代码背后的思路值得多说两句。第一种子设成42是为了保证每次随机出的数据一样方便复现。第二噪声系数2决定了数据集的难度——噪声调小点几乎落在直线上模型好拟合同时也学不到什么噪声调大点的分布很散回归线的斜率会被噪声拉扯。做实验时我习惯先设一个中等噪声等代码跑通了再逐步加大观察模型抗干扰能力。第三如果把np.random.rand(100,1)改成np.random.rand(100,2)就从一元线性回归变成了多元线性回归代码逻辑几乎不变但涉及的概念特征维度、系数向量会复杂一截。这篇先牢牢盯住一元。画出散点图确认一下数据结构这是流程里不能省的一步plt.scatter(X, y, alpha0.6) plt.xlabel(X) plt.ylabel(y) plt.title(模拟数据分布) plt.show()2.3 最小二乘法在代码里到底在算什么你可能看过最小二乘的公式要求的就是让每个点的预测值和真实值差的平方和最小。代码里的实现逻辑和公式推导是一致的但有一个关键区别推导时用矩阵形式写起来干净代码实现时更常直接写循环或向量化运算。具体到损失函数def compute_loss(y_true, y_pred): return np.mean((y_true - y_pred) ** 2)这个函数返回的就是均方误差MSE。np.mean而不是np.sum好处是把误差归一化到每个样本平均误差的量级这样不管数据集有100条还是10000条记录Loss的数值范围都会比较稳定方便你判断模型是否收敛。3. 最简单的回归代码两种写法手写与调库3.1 手写梯度下降版二十几行看清回归的本质梯度下降的思路用生活场景类比就是你站在山坡上想走到谷底每次左右看看哪个方向是下坡迈出一步再重复——步子太大可能跳过谷底步子太小走得太慢。代码里的方向就是梯度步子就是学习率。下面是我常给学生演示的最小实现核心就一个循环# 初始化参数 w, b 0.0, 0.0 learning_rate 0.01 epochs 1000 n len(X) for epoch in range(epochs): # 1. 预测 y_pred w * X.squeeze() b # 2. 计算损失 loss np.mean((y_pred - y) ** 2) # 3. 计算梯度对w求偏导、对b求偏导 dw (2 / n) * np.dot(X.squeeze(), y_pred - y) db (2 / n) * np.sum(y_pred - y) # 4. 更新参数 w - learning_rate * dw b - learning_rate * db if epoch % 100 0: print(fEpoch {epoch}: loss{loss:.4f}, w{w:.2f}, b{b:.2f}) print(f最终结果: w{w:.2f}, b{b:.2f}真实值: w3.00, b5.00)代码理解拆成三层第一层预测公式w * X b对应线性回归的假设函数这是所有代码的地基。第二层梯度计算里的(2/n)来自均方误差对w和b求导的结果如果你推过公式会发现手写代码没有绕过数学只是把求导结果直接落地了。第三层w - learning_rate * dw是梯度下降的更新规则负号表示往损失减小的方向走。跑完这段代码如果一切正常w会在3.0附近b在5.0附近——因为造数据时我们就是按y 3x 5生成的。这正是手写代码最大的好处用已知答案的数据验证实现逻辑。如果跑出来w是3.5甚至4.0说明梯度下降过程有问题而不是模型有问题。3.2 三段式调库scikit-learn的正式用法手写版跑通后就该看生产环境下真正常用的写法了。实际项目中几乎没人手写梯度下降都用现成库。用scikit-learn完成线性回归核心代码就三行式的流程from sklearn.linear_model import LinearRegression from sklearn.model_selection import train_test_split from sklearn.metrics import mean_squared_error, r2_score # 1. 划分训练集和测试集 X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.2, random_state42 ) # 2. 训练 model LinearRegression() model.fit(X_train, y_train) # 3. 预测并评估 y_pred model.predict(X_test) mse mean_squared_error(y_test, y_pred) r2 r2_score(y_test, y_pred) print(f系数w: {model.coef_[0]:.2f}) print(f截距b: {model.intercept_:.2f}) print(f测试集MSE: {mse:.4f}) print(f测试集R²: {r2:.4f})有个细节特别容易让新手困惑为什么手写版直接用了全部数据调库版却要先划分训练集和测试集原因在于评估方式的不同。手写版我们是在验证实现逻辑——数据是合成的、答案已知不需要担心过拟合调库版要模拟真实场景——模型得在没见过的数据上表现好才算有用如果只用训练数据评估哪怕模型完全记住数据点也能拿到很低的MSE但这没有任何预测价值。train_test_split里test_size0.2意味着80%数据训练、20%数据测试这是最常见的一种划分比例。3.3 两种结果对拍用同一份数据相互验证手写版和调库版之间应该是什么关系答案是结果应该高度接近但不要求完全相等。手写版是梯度下降迭代求近似解迭代1000次后收敛到的位置接近全局最优但可能有微小误差调库版用的是最小二乘的闭式解直接解方程得到精确解。以我的经验学习率合理、迭代次数足够时两者差距通常在0.01以内。如果差距大先查学习率再看迭代次数最后检查是不是数据划分不一致。对拍验证的脚本我长期保留换个数据集就能用# 把两个模型的结果放在同一张图上看 plt.scatter(X_test, y_test, alpha0.6, label真实测试点) plt.plot(X_test, y_pred, r-, labelsklearn回归线) # 手写模型的系数也画一条线 plt.plot(X_test, manual_w * X_test.squeeze() manual_b, g--, label手写梯度下降回归线) plt.legend() plt.show()两张线如果几乎重合说明你既理解了原理也会用工具这一章就算真正过关了。4. 实操中的五个典型问题与排查思路4.1 数据没做归一化导致梯度震荡如果数据里某个特征取值范围是0到100000另一个特征是0到1梯度下降就会出问题大数值特征的梯度幅值很大小数值特征的梯度很小两者相差几个数量级更新参数时要么大特征方向步子过猛要么小特征方向基本不动训练过程来回震荡。我的排查习惯一旦发现loss曲线像锯齿一样上下颠簸先打印出w和b的变化值看是不是某一个方向更新量特别大。解决办法很直接标准化或归一化后重新跑from sklearn.preprocessing import StandardScaler scaler StandardScaler() X_scaled scaler.fit_transform(X)注意一点fit_transform之后做预测新数据也要用同一个scaler做transform不能用原来的原始值往模型里塞否则特征分布对不上预测结果直接飘掉。这个问题在博文评论区经常有人问说是练得好好的一用真实数据就乱套多半就是漏了这一步。4.2 学习率设得不对NaN和发散学习率太大参数更新直接越过最低点甚至每一步都跳得比上一步更远loss一路飙升变成inf最终出现NaN。学习率太小loss下降非常缓慢1000次迭代后参数离最优值还差很远。每次换数据集都要把学习率当作要重新调整的超参数来看不要沿用上次的默认值。我调试时的经验方法先把学习率设成0.001跑一遍观察loss是否稳定下降。如果下降太慢改成0.01再看如果发散了改0.0001。用这种十倍试错的方式快速定位一个量级合适的区间。还有个偷懒技巧给loss更新设一个打印阈值连续50次迭代loss下降幅度小于1e-5就提前停止省时间还避免过拟合到迭代次数上if epoch 0 and abs(loss - prev_loss) 1e-5: print(f提前收敛停止于 epoch {epoch}) break prev_loss loss4.3 数据里有NaN或无穷值这个坑的隐蔽之处在于大部分情况代码不会直接报错而是默默返回一个nan或者离谱的系数。你看到loss变成nan的第一反应通常是梯度爆炸了但先别急着调学习率——先检查原始数据import pandas as pd df pd.DataFrame({X: X.squeeze(), y: y}) print(df.isna().sum()) print(np.isinf(df).sum())如果确有缺失常用策略是删除或填充。小数据集直接删除缺失行大数据集用均值/中位数填充更稳妥。fillna的mean和median在sklearn的SimpleImputer里也有封装字段名差不多别记混。这个排查思路对后面学任何模型都有用因为数据清洗永远是第一步。4.4 R²MSE和感觉不对的判断回归模型有两个最常用的评估指标MSE和R²。MSE衡量的是预测误差的平均水平数值越小越好但它有量纲不好跨数据集比较R²是无量纲的取值范围通常在0到1之间可以通俗理解成模型解释了百分之多少的方差0.8意味着y的变化有八成被模型抓住了。但这两套指标在真实数据上都有失灵的时候。如果测试集R²是负的说明你的模型预测水平还不如直接拿y的均值去猜基本可以断定模型结构有问题或数据预处理有问题。另外别只看单一指标我遇到过一次MSE挺小但画图一看拟合线是歪的原因是数据里有几个极端异常点把回归线拽偏了MSE被多数正常点拉低视觉上却明显不对。所以我的习惯永远是指标看图也画两者参照着下结论。4.5 一元回归没问题但多元回归跑崩了系列的下一篇可能就会涉及多元回归这里先打一个预防针一元回归里你能直接画散点图看线性关系多元回归的维度超过3就没法直接可视化这时候更依赖指标和诊断图。同时多元回归最容易出问题的点是多重共线性——两个特征高度相关时系数估计会非常不稳定甚至符号都反常识。判断方法是看相关系数矩阵相关系数超过0.8就要警惕。处理办法是删掉一个冗余特征改用岭回归或Lasso这个内容在下一节展开。5. 从最简单起步线性回归之后的扩展方向5.1 正则化面对不对劲的线性回归真实业务数据很少像模拟数据那么规整特征多、相关性高、噪声大普通线性回归容易过拟合或系数解释性差。这时候岭回归L2正则和LassoL1正则是线性回归最自然的升级。正则化的原理一句话就能说清在损失函数后面加一个对系数大小的惩罚项逼着模型在拟合得好和系数别太夸张之间取平衡。代码上只是换了个模型名字from sklearn.linear_model import Ridge, Lasso ridge Ridge(alpha1.0) ridge.fit(X_train, y_train) lasso Lasso(alpha0.01) lasso.fit(X_train, y_train)alpha越大惩罚越强系数被压得越小但它不是越大越好——压得太狠模型就欠拟合了。alpha的调参需要用交叉验证GridSearchCV或者RidgeCV这种带CV后缀的类不能用测试集反复试否则会信息泄露后续评估虚高。5.2 非线性怎么办从线性回归到树模型线性回归的局限性非常明确它假设特征和标签之间是直线关系。实际场景里更多是非线性关系比如年龄和收入的关系、温度和销量的关系用线性模型硬拟合的效果很差。这时候有两个方向方向一对特征做变换比如加平方项、交互项把非线性关系掰成线性后再套线性回归。优点是可解释性好缺点是变换公式需要人工设计和业务经验。方向二换模型比如决策树回归、随机森林回归。我现在遇到非线性关系明显的数据第一反应就是直接上树模型因为它不需要做特征缩放、能捕捉复杂交互关系写代码也不麻烦from sklearn.ensemble import RandomForestRegressor rf RandomForestRegressor(n_estimators200, max_depth10, random_state42) rf.fit(X_train, y_train)随机森林里有几个关键参数值得记住n_estimators是树的数量一般100到300之间效果差异不大但训练时间会线性增加max_depth控制单棵树的深度太深容易过拟合训练集random_state固定之后保证结果可复现。关于随机森林的调参建议是优先调max_depth和min_samples_leaf而不是无脑加树的数量。再往上走就是梯度提升树XGBoost和LightGBM。这两个模型在工业界用得非常多原因是它们把多个弱学习器逐步提升的思路落地成工程化工具速度更快、精度更高。以LightGBM为例一份可以跑的基础代码import lightgbm as lgb model lgb.LGBMRegressor( n_estimators500, learning_rate0.05, max_depth-1, num_leaves31, random_state42, ) model.fit(X_train, y_train)LightGBM几个入门的参数常识learning_rate决定了每棵树的贡献权重调小以后需要更多树来弥补num_leaves是核心复杂度控制参数默认31在中小数据集上够用max_depth设-1表示不限制但配合num_leaves控制复杂度。建议先跑默认参数再看特征重要性做筛选最后调参不要一开始就追求最优。5.3 学习路径建议从这篇笔记往后怎么走以我自己的项目经验来看回归模型的学习路径可以画成一条清晰的线线性回归本文搞懂损失、梯度下降、评估指标这是后续所有模型的地基。正则化线性模型岭回归/Lasso解决特征冗余和数据噪声问题。决策树回归 → 随机森林引入非线性能力同时建立集成学习的直觉。XGBoost/LightGBM工业界的默认选择特征工程做好之后直接出成绩。回看线性回归的假设诊断残差图、Q-Q图、共线性诊断让自己能从会跑代码进阶到能判断模型是否可靠。这里必须强调一个容易被忽略的点这套路径的价值不在于每个模型都学会调用API而在于同一个数据集上做横向对比。同一个任务用线性回归、岭回归、随机森林、LightGBM各跑一遍记录各自的MSE、R²、训练时间你才能直观感受什么场景该用什么模型。我每次带项目都会让组员做成一张对比表这张表比任何理论讲解都更有冲击力。6. 把线性回归写进自己的工具箱这篇文章从liner这个拼写开头一路写到岭回归、LightGBM核心其实就是一个观点线性回归不是一学完就可以丢掉的玩具而是判断一切回归问题的起点。我整理一下自己常用的代码片段你可以直接当模板存一份# 1. 手写梯度下降理解原理时用 # 2. sklearn LinearRegression基线模型永远先跑这个 # 3. Ridge / Lasso遇到共线性或过拟合时替换 # 4. RandomForestRegressor / LGBMRegressor非线性问题或追求精度时上每次拿到一份新数据集我拿到手的第一反应永远是先跑线性回归做baseline。它不是效果最好的模型但它能告诉我很多信息特征和标签的大致关系、数据质量是否靠谱、评估流程是否通顺。等这些底都摸清了再决定要不要上复杂模型心里就有底了。最后分享两个个人习惯。第一模拟数据是调试代码最好的朋友先造已知答案的数据跑出来的系数对不对一目了然等逻辑确认无误再换真实数据能帮你省下一大半为什么结果这么离谱的排查时间。第二写回归代码时始终带着一个疑问这个预测结果我能不能用一两句话说清楚它为什么是根据这些特征得出这个值——如果连自己都解释不了那这个模型多半还没调到位。希望这篇笔记能给你的回归学习之旅省点力气。代码量不大重点是动手敲一遍改几个参数看看会发生什么。
返回列表