别再死记硬背了!用Python手动画出PyTorch线性回归的损失函数曲线(附完整代码)
用Python可视化PyTorch线性回归:从数学公式到损失函数曲线的直觉理解
当你第一次接触线性回归时,那些数学公式可能看起来像天书一样抽象。为什么损失函数是碗状的?最优权重w到底在哪里?这些问题困扰着许多初学者。本文将带你用Python和Matplotlib,通过可视化手段把这些抽象概念变成直观的图形,让你真正"看到"机器学习背后的数学原理。
1. 准备工作与环境搭建
在开始之前,我们需要确保环境配置正确。这个项目只需要几个基础的Python库:
import numpy as np
import matplotlib.pyplot as plt
为什么选择这两个库? NumPy提供了高效的数值计算能力,而Matplotlib则是Python中最常用的可视化工具。它们组合起来,可以让我们轻松实现从计算到可视化的完整流程。
对于这个示例,我们使用一个简单的数据集:
x_data = [1.0, 2.0, 3.0] # 输入特征
y_data = [2.0, 4.0, 6.0] # 对应的标签
这个数据集描述了一个完美的线性关系:y = 2x。虽然简单,但足够展示线性回归的核心概念。
2. 理解线性模型与损失函数
线性回归模型的基本形式是:
ŷ = w * x
其中:
- ŷ是我们的预测值
- w是权重参数(也就是我们要找的"最佳斜率")
- x是输入特征
损失函数则衡量了预测值与真实值之间的差距。最常用的均方误差(MSE)定义为:
MSE = 1/N * Σ(ŷ - y)²
让我们用Python实现这两个关键函数:
def forward(x):
"""前向传播函数,计算预测值ŷ"""
return x * w
def loss(x, y):
"""计算单个数据点的损失"""
y_pred = forward(x)
return (y_pred - y) ** 2
注意:这里的w是一个全局变量,我们将在后面的循环中尝试不同的w值。
3. 计算并可视化损失函数曲线
现在到了最有趣的部分——绘制损失函数曲线。我们将:
- 尝试一系列不同的w值(从0.0到4.0,步长0.1)
- 对每个w值,计算整个数据集的平均损失(MSE)
- 记录这些w值和对应的MSE值
- 最后绘制w-MSE曲线
实现代码如下:
w_list = [] # 存储不同的w值
mse_list = [] # 存储对应的MSE值
for w in np.arange(0.0, 4.0, 0.1): # w从0.0到4.0,步长0.1
l_sum = 0 # 初始化损失总和
for x_val, y_val in zip(x_data, y_data):
l = loss(x_val, y_val) # 计算单个数据点的损失
l_sum += l # 累加损失
mse = l_sum / len(x_data) # 计算平均损失(MSE)
w_list.append(w)
mse_list.append(mse)
print(f"w={w:.1f}, MSE={mse:.2f}") # 打印当前w和MSE
运行这段代码后,我们可以绘制损失函数曲线:
plt.plot(w_list, mse_list)
plt.xlabel('Weight (w)')
plt.ylabel('Mean Squared Error (MSE)')
plt.title('Loss Function for Linear Regression')
plt.show()
你会看到一个漂亮的"碗状"曲线,这就是我们要找的损失函数可视化结果!
4. 深入分析损失函数曲线
观察我们得到的曲线,有几个关键点值得注意:
-
曲线形状:为什么是碗状(凸函数)?
- 当w偏离最优值(在这个例子中是2.0)时,MSE会对称地增加
- 这种对称性来自于平方误差的性质
-
最小值点:
- 曲线的最低点对应着最优的w值
- 在这个例子中,最小值出现在w=2.0处
- 这与我们的数据生成规律y=2x完美吻合
-
梯度下降的直观理解:
- 想象一个小球从w=0.0处开始滚动
- 它会自然地沿着曲线下滑,最终停在最低点
- 这就是梯度下降算法的直观表现
为了更清楚地看到这一点,我们可以标记出最小值点:
min_index = np.argmin(mse_list)
min_w = w_list[min_index]
min_mse = mse_list[min_index]
plt.plot(w_list, mse_list)
plt.scatter(min_w, min_mse, color='red', label=f'Minimum (w={min_w:.1f}, MSE={min_mse:.2f})')
plt.xlabel('Weight (w)')
plt.ylabel('Mean Squared Error (MSE)')
plt.title('Loss Function with Minimum Point Marked')
plt.legend()
plt.show()
5. 扩展实验与思考
理解了基本概念后,我们可以进行一些有趣的扩展实验:
实验1:改变数据分布
尝试修改y_data,看看曲线如何变化:
y_data = [2.1, 3.9, 6.2] # 添加一些噪声
实验2:尝试不同的损失函数
比如绝对误差(MAE):
def loss_mae(x, y):
y_pred = forward(x)
return abs(y_pred - y)
实验3:多参数情况
对于y = w1 * x + w0的模型,损失函数将是一个三维曲面。我们可以用类似的方法可视化:
from mpl_toolkits.mplot3d import Axes3D
w0_range = np.arange(-2, 2, 0.1)
w1_range = np.arange(0, 4, 0.1)
W0, W1 = np.meshgrid(w0_range, w1_range)
MSE = np.zeros_like(W0)
for i in range(len(w0_range)):
for j in range(len(w1_range)):
w0 = w0_range[i]
w1 = w1_range[j]
l_sum = 0
for x_val, y_val in zip(x_data, y_data):
y_pred = w1 * x_val + w0
l = (y_pred - y_val) ** 2
l_sum += l
MSE[j, i] = l_sum / len(x_data)
fig = plt.figure()
ax = fig.add_subplot(111, projection='3d')
ax.plot_surface(W0, W1, MSE, cmap='viridis')
ax.set_xlabel('w0 (bias)')
ax.set_ylabel('w1 (weight)')
ax.set_zlabel('MSE')
plt.show()
通过这些实验,你会对线性回归和损失函数有更深入的理解。可视化不仅帮助我们理解算法,还能在调试模型时提供直观的反馈。
更多推荐


所有评论(0)