从零实现Linear Regression报错:我的代码为何无法正常运行?
线性回归梯度下降实现中的数值溢出与NaN问题
我尝试用Python从零实现线性回归,参考最小二乘损失和梯度下降的数学公式。
实现代码
import numpy as np class LinearRegression: def __init__( self, features: np.ndarray[np.float64], targets: np.ndarray[np.float64], ) -> None: self.features = np.concatenate((np.ones((features.shape[0], 1)), features), axis=1) self.targets = targets self.params = np.random.randn(features.shape[1] + 1) self.num_samples = features.shape[0] self.num_feats = features.shape[1] self.costs = [] def hypothesis(self) -> np.ndarray[np.float64]: return np.dot(self.features, self.params) def cost_function(self) -> np.float64: pred_vals = self.hypothesis() return (1 / (2 * self.num_samples)) * np.dot((pred_vals - self.targets).T, pred_vals - self.targets) def update(self, alpha: np.float64) -> None: self.params = self.params - (alpha / self.num_samples) * (self.features.T @ (self.hypothesis() - self.targets)) def gradientDescent(self, alpha: np.float64, threshold: np.float64, max_iter: int) -> None: converged = False counter = 0 while not converged: counter += 1 curr_cost = self.cost_function() self.costs.append(curr_cost) self.update(alpha) new_cost = self.cost_function() if abs(new_cost - curr_cost) < threshold: converged = True if counter > max_iter: converged = True
调用方式
regr = LinearRegression(features=np.linspace(0, 1000, 200, dtype=np.float64).reshape((20, 10)), targets=np.linspace(0, 200, 20, dtype=np.float64)) regr.gradientDescent(0.1, 1e-3, 1e+3) regr.cost_function()
报错信息
RuntimeWarning: overflow encountered in scalar power return (1 / (2 * self.num_samples)) * (la.norm(self.hypothesis() - self.targets) ** 4) RuntimeWarning: invalid value encountered in scalar subtract if abs(new_cost - curr_cost) < threshold: RuntimeWarning: overflow encountered in matmul self.params = self.params - (alpha / self.num_samples) * (self.features.T @ (self.hypothesis() - self.targets))
问题原因分析
- 特征尺度未归一化/标准化:特征取值范围是0-1000,目标值仅为0-200,特征尺度远大于目标。梯度下降中参数更新量和特征值直接相关,大尺度特征会导致梯度值异常大,参数迅速溢出为
inf,后续计算就会出现NaN。 - 学习率α过大:在特征未缩放的情况下,设置α=0.1会让参数更新步幅远超合理范围,不仅无法收敛,还会直接引发数值爆炸。
- 损失函数实现不一致:报错信息中的损失函数使用了四次方(
**4),但贴出的代码是平方和。如果实际运行的是四次方版本,损失值增长速度会更快,进一步加剧溢出问题。
解决办法
- 特征缩放:对输入特征做Z-score标准化(减去均值除以标准差)或Min-Max归一化(缩放到0-1区间),确保所有特征处于相近尺度。
- 降低学习率:配合特征缩放,将α调整为0.001或更小,逐步调试找到合适的取值。
- 修正损失函数:确保损失函数实现和参考公式一致,回归任务中常规使用均方误差的一半(平方和形式),而非四次方。
内容的提问来源于stack exchange,提问作者Sagnik Taraphdar
相关产品推荐
相关产品推荐

