批量梯度下降实现线性Regression遇MSE不下降问题求助
批量梯度下降线性回归模型MSE不下降问题排查
我实现了一个基于批量梯度下降的线性回归模型,但训练过程中均方误差(MSE)始终没有下降。LinearModel是一个预定义模板类,初始化超参数为step_size=0.001、max_iter=10000、eps=0.001、theta_0=None、verbose=True,使用的是Fish Market数据集。
模型代码
# 数据存储为(m, n)形状的数组,n列m行 # x = array((m,n)) # y = array((m,1)) class LinearRegression(LinearModel): """基于梯度下降的线性回归""" def fit(self, x, y): """运行梯度下降法最小化线性回归的损失函数J(theta) :param x: 训练样本输入,形状(m, n) :param y: 训练样本标签,形状(m,) """ def squared_error(theta, x,y): # 计算均方误差 return np.mean((x.dot(theta) - y)**2) # 计算均方误差对每个参数的偏导数 def grad_squared_error(theta, x,y): m,n = x.shape grad = np.zeros(n) for j in range(n): # 用链式法则更新偏导数 grad[j] += ((x.dot(theta)) - y).dot(x[:,j][np.newaxis].T) # np.newaxis将数组转为2D以便转置 # 返回偏导数除以(数据量*2) return [g / (2*m) for g in grad] m,n = x.shape if self.theta is None: self.theta = np.zeros(n) decay_factor = 0.9 # 可调整衰减系数 decay_interval = 10 # 可调整衰减间隔 for i in range(self.max_iter): # 计算当前参数下的损失偏导数 grad = grad_squared_error(self.theta, x,y) # 逐个更新参数 for j in range(len(self.theta)): self.theta[j] -= self.step_size * grad[j] # 打印当前参数和误差并衰减学习率 if self.verbose and i % decay_interval == 0: self.step_size *= decay_factor print(f"Iteration {i+1}: theta = {self.theta}, squared error = {squared_error(self.theta, x,y)}") if squared_error(self.theta, x,y) < self.eps: if self.verbose: print("已收敛,参数如下:") print(f"Iteration {i+1}: theta = {self.theta}, squared error = {squared_error(self.theta, x,y)}") break def predict(self, x): """对新输入x做出预测 :param x: 输入数据,形状(m, n) :return: 预测输出,形状(m,) """ return np.dot(x, self.theta)
迭代后期输出
Iteration 9911: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9921: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9931: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9941: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9951: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9961: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9971: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9981: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878 Iteration 9991: theta = [-0.47145182 5.37944606 6.10798726 5.03442352 4.33745377 1.48791336], squared error = 19565.226327190878
我尝试过每10步衰减学习率,但问题仍未解决。
问题根源及修正方案
1. 梯度计算错误
线性回归MSE损失的梯度公式应为$\frac{1}{m}X^T(X\theta - y)$,你的代码存在两个问题:
- 梯度最后除以了
2*m,但你的squared_error用的是np.mean((x.dot(theta)-y)**2)(对应$\frac{1}{m}\sum(...)2$),梯度应该是$\frac{2}{m}XT(X\theta - y)$,或者统一损失和梯度的计算逻辑。 - 用列表推导式把numpy数组转成普通列表,后续参数更新效率低且易出错,应保持numpy数组类型。
修正后的梯度计算函数:
def grad_squared_error(theta, x, y): m, n = x.shape error = x.dot(theta) - y # 向量化计算,无需循环 grad = x.T.dot(error) / m return grad
2. 学习率衰减时机错误
你在第一次迭代(i=0)就开始衰减学习率,导致学习率迅速极小化,参数几乎无法更新。调整衰减时机,跳过首次迭代:
for i in range(self.max_iter): grad = grad_squared_error(self.theta, x, y) # 向量化更新参数,替代循环 self.theta -= self.step_size * grad if self.verbose and i % decay_interval == 0: print(f"Iteration {i+1}: theta = {self.theta}, squared error = {squared_error(self.theta, x,y)}") # 每10步衰减一次,跳过第一次迭代 if i > 0 and i % decay_interval == 0: self.step_size *= decay_factor current_error = squared_error(self.theta, x, y) if current_error < self.eps: if self.verbose: print("已收敛,参数如下:") print(f"Iteration {i+1}: theta = {self.theta}, squared error = {current_error}") break
3. 数据未做标准化
Fish Market数据集中的特征数值范围差异极大(比如体重是几百,长度是几十),梯度下降对未标准化的数据非常敏感,会导致梯度更新失衡,无法有效收敛。必须对输入特征做标准化处理:
from sklearn.preprocessing import StandardScaler # 标准化特征 scaler = StandardScaler() x_scaled = scaler.fit_transform(x) # 使用标准化后的数据训练模型 model = LinearRegression() model.fit(x_scaled, y)
4. 参数更新向量化
原代码用for循环逐个更新theta,效率低且易出错,改用numpy向量化操作可避免这些问题,同时提升运行速度。
内容的提问来源于stack exchange,提问作者Hassan Abbas
相关产品推荐
相关产品推荐

