线性回归梯度下降:能否绘制待最小化的目标函数?
当然可以!梯度下降要最小化的**损失函数(对线性回归来说通常是均方误差MSE)**完全能可视化,而且这绝对是帮你理解梯度下降工作逻辑的好办法——你能亲眼看到算法是怎么一步步朝着最小值“挪”过去的。
一、不同场景下的可视化方式
根据你线性回归模型的特征数量,可视化的形式会略有不同:
- 单特征模型:损失函数是关于截距
θ₀和斜率θ₁的二元函数,可以画3D曲面,或者固定其中一个参数,画另一个参数与损失的2D抛物线(开口向上,最低点就是最优参数)。 - 双特征模型(就像你代码里的二维输入):损失函数是关于
θ₀、θ₁、θ₂的三元函数,我们可以固定其中一个参数(比如固定θ₀),然后画θ₁、θ₂与损失的3D碗状曲面,或者等高线图,直观看到最小值的位置。
二、基于你数据的可视化代码示例
假设你用的是线性回归最常用的均方误差损失函数,我基于你给出的X数据补全了可视化代码(你只需要替换成自己的真实标签y即可):
import matplotlib.pyplot as plt import numpy as np # 你的输入特征数据(补全了一条示例数据,你可以替换成完整数据集) X = np.array([ [2.13, 5.49], [8.35, 6.74], [8.17, 5.79], [0.62, 8.54], [2.34, 4.87] ]) # 替换成你自己的真实标签数据 y = np.array([12.3, 20.1, 18.9, 10.5, 11.7]) # 定义均方误差损失函数 def compute_mse_loss(theta0, theta1, theta2, X, y): sample_count = len(y) # 线性回归的预测值 y_pred = theta0 + theta1 * X[:, 0] + theta2 * X[:, 1] # 计算损失(1/(2m)是为了求导时简化计算,不影响最小值位置) loss = (1/(2 * sample_count)) * np.sum((y_pred - y)**2) return loss # 生成参数网格:固定theta0为0,遍历theta1和theta2的可能取值 theta0_fixed = 0 theta1_range = np.linspace(-5, 10, 100) theta2_range = np.linspace(-5, 10, 100) theta1_grid, theta2_grid = np.meshgrid(theta1_range, theta2_range) # 计算每个参数组合对应的损失值 loss_grid = np.zeros_like(theta1_grid) for i in range(theta1_grid.shape[0]): for j in range(theta1_grid.shape[1]): loss_grid[i,j] = compute_mse_loss(theta0_fixed, theta1_grid[i,j], theta2_grid[i,j], X, y) # 绘制3D损失曲面和等高线图 fig = plt.figure(figsize=(12, 6)) # 3D曲面图 ax1 = fig.add_subplot(121, projection='3d') ax1.plot_surface(theta1_grid, theta2_grid, loss_grid, cmap='viridis', alpha=0.8) ax1.set_xlabel('θ₁(第一个特征的系数)') ax1.set_ylabel('θ₂(第二个特征的系数)') ax1.set_zlabel('损失值 J(θ)') ax1.set_title('3D损失曲面(固定θ₀=0)') # 等高线图 ax2 = fig.add_subplot(122) contour_plot = ax2.contour(theta1_grid, theta2_grid, loss_grid, levels=20, cmap='viridis') ax2.clabel(contour_plot, inline=True, fontsize=8) ax2.set_xlabel('θ₁') ax2.set_ylabel('θ₂') ax2.set_title('损失函数等高线图') plt.tight_layout() plt.show()
三、可视化后的直观结论
运行这段代码后,你会看到:
- 3D图里是一个光滑的碗状曲面,碗底就是损失最小的参数组合;
- 等高线图里的闭合曲线代表相同损失的参数集合,越靠近中心的曲线损失越小,中心就是最优解。
另外补充一句:如果你代码里的sigmoid是误写(毕竟你说的是线性回归),如果是逻辑回归的交叉熵损失函数,可视化逻辑是一样的——交叉熵也是凸函数,形态类似碗状,同样可以用上面的思路绘制。
内容的提问来源于stack exchange,提问作者lukassz
相关产品推荐
相关产品推荐

