这段代码是一个使用 PyTorch 从零实现线性回归(Linear Regression)的完整示例。它没有使用 torch.nn 中的高级封装(如 nn.Linear 或 nn.MSELoss),而是手动定义了模型、损失函数和优化器,旨在帮助理解深度学习框架底层的运作机制(如自动求导、梯度下降)。
下面我将分模块详细解析这段代码:
1. 导入库
import torch
import matplotlib.pyplot as plt # 画图的
import random
torch: PyTorch 核心库,用于张量运算和自动求导。matplotlib.pyplot: 用于数据可视化。random: 用于打乱数据索引。
2. 数据生成 (create_data)
def create_data(w,b,data_num):
x = torch.normal(0,1,(data_num,len(w))) # 生成标准正态分布的特征 X
y = torch.matmul(x,w) + b # 计算真实的标签 y = Xw + b
noise = torch.normal(0,0.01,y.shape) # 生成噪声
y += noise # 给标签加入噪声,模拟真实数据
return x,y
- 功能: 构造 synthetic data(合成数据)。
- 逻辑: 假设真实关系是线性的。
x是输入特征(500 行,4 列),w是权重,b是偏置。加入高斯噪声noise是为了让数据不那么完美,模拟现实世界中的误差,测试模型的拟合能力。
3. 数据准备与可视化
num=500
true_w = torch.tensor([8.1,2,2,4])
true_b = torch.tensor(1.1)
X,Y = create_data(true_w,true_b,num)
plt.scatter(X[:,0],Y,1) # 只画第一个特征和标签的关系
plt.show()
- 设置了真实的权重
true_w和偏置true_b。 - 生成了 500 条数据。
- 注意: 这里只画了
X的第一列特征与Y的散点图,因为 4 维数据无法直接在 2D 平面完全展示。
4. 数据迭代器 (data_provider)
def data_provider(data,label,batchsize):
length = len(label)
indices = list(range(length))
random.shuffle(indices) # 打乱索引,防止模型记忆数据顺序
for each in range(0,length,batchsize):
get_indices = indices[each:each+batchsize]
get_data = data[get_indices]
get_label = label[get_indices]
yield get_data,get_label # 生成器,每次返回一个 batch
- 功能: 实现小批量随机梯度下降 (Mini-batch SGD) 所需的数据加载。
- 关键点:
random.shuffle: 每个 epoch 打乱数据,增加随机性,有助于收敛。yield: 使用 Python 生成器,节省内存,不用一次性把所有 batch 都存下来。batchsize = 16: 每次训练只用 16 个样本计算梯度。
5. 模型定义 (fun)
def fun(x,w,b):
pred_y = torch.matmul(x,w) + b
return pred_y
- 功能: 定义前向传播(Forward Pass)。
- 公式: y^=Xw+b 。这里手动实现了矩阵乘法。
6. 损失函数 (maeLoss)
def maeLoss(pre_y,y):
return torch.sum(abs(pre_y-y))/len(y)
- 功能: 计算预测值与真实值之间的误差。
- 类型: 这里使用的是 MAE (Mean Absolute Error, 平均绝对误差),即 L1 Loss。
- 公式:n1∑∣y^−y∣ 。
- 注:回归问题中更常用 MSE (均方误差),但 MAE 也是合法的损失函数。
7. 优化器 (sgd)
def sgd(paras,lr):
with torch.no_grad(): # 更新参数时不需要记录梯度,节省内存
for para in paras:
para -= para.grad * lr # 梯度下降公式:w = w - lr * grad
para.grad.zero_() # 重要:清空梯度,否则梯度会累加
- 功能: 手动实现随机梯度下降更新规则。
- 关键点:
torch.no_grad(): 参数更新本身不需要构建计算图。para.grad.zero_(): 非常重要。PyTorch 默认会累加梯度(为了 RNN 等场景),所以在每次更新后必须手动清零,否则下一次backward()会出错或结果不对。
8. 初始化参数
lr = 0.03
# 初始化权重,requires_grad=True 表示需要计算该张量的梯度
w_0 = torch.normal(0,0.01,true_w.shape,requires_grad=True)
b_0 = torch.tensor(0.01,requires_grad=True)
requires_grad=True: 这是 PyTorch 自动求导的开关。只有设置为 True 的张量,PyTorch 才会追踪其操作并计算梯度。- 参数初始化为接近 0 的随机数。
9. 训练循环
epochs = 50
for epoch in range(epochs):
data_loss = 0
for batch_x,batch_y in data_provider(X,Y,batchsize):
pred_y = fun(batch_x,w_0,b_0) # 1. 前向传播
loss = maeLoss(pred_y,batch_y) # 2. 计算损失
loss.backward() # 3. 反向传播 (自动计算梯度)
sgd([w_0,b_0],lr) # 4. 更新参数
data_loss += loss
print("epoch %03d: loss: %.6f"%(epoch,data_loss))
- 流程:
- Forward: 计算预测值。
- Loss: 计算标量损失。
- Backward:
loss.backward()会利用链式法则,自动计算w_0和b_0的梯度,并存入.grad属性中。 - Step: 调用自定义的
sgd函数更新权重并清零梯度。
- 输出: 每个 epoch 打印一次总损失,理论上损失应该随 epoch 增加而下降。
10. 结果验证与绘图
print("函数值",true_w,true_b)
print("训练值",w_0,b_0)
idx = 0
# 绘制拟合直线
plt.plot(X[:, idx].detach().numpy(), X[:,idx].detach().numpy()*w_0[idx].detach().numpy()+b_0.detach().numpy())
plt.scatter(X[:, idx],Y,1)
plt.show()
- 打印真实参数和训练后的参数,对比看看是否接近。
- 绘图潜在问题:
plt.plot会按照X[:, idx]在数组中的顺序连线。- 由于
X是随机生成的,X[:, 0]并不是从小到大排序的。 - 后果: 画出来的拟合线会是杂乱的折线(像一团毛线),而不是一条直的回归线。
- 修正建议: 绘图前应该先对
X[:, idx]进行排序,或者使用散点图展示预测值。
代码核心知识点总结
- 计算图 (Computational Graph): PyTorch 通过追踪对
requires_grad=True的张量的操作来构建动态计算图。 - 自动求导 (Autograd):
loss.backward()是核心,它自动计算损失函数对所有叶子节点(这里是w_0,b_0)的梯度。 - 梯度清零: 在 PyTorch 中,梯度是累加的,每次迭代后必须
zero_()。 - Detach: 绘图时使用
.detach().numpy(),将张量从计算图中分离并转为 NumPy 数组,因为 matplotlib 不处理 PyTorch 张量,且绘图不需要梯度信息。
改进建议 (Best Practices)
虽然这段代码教学意义很好,但在实际工程中,我们通常这样做:
- 使用
nn.Module: 继承torch.nn.Module定义模型,管理参数更方便。 - 使用
optim: 使用torch.optim.SGD代替手写的sgd函数。 - 使用
DataLoader: 使用torch.utils.data.DataLoader代替手写的data_provider,支持多进程加载。 - 修复绘图:
# 修复绘图乱线的问题 x_plot = X[:, idx].detach().numpy() y_plot = x_plot * w_0[idx].detach().numpy() + b_0.detach().numpy() # 按 x 排序 sorted_indices = np.argsort(x_plot) plt.plot(x_plot[sorted_indices], y_plot[sorted_indices], 'r-') - 损失函数: 回归问题通常首选
MSELoss(均方误差),因为 MAE 在 0 点不可导(虽然 PyTorch 能处理次梯度),且 MSE 对大误差惩罚更重,收敛通常更稳定。
总体来说,这是一段非常经典的**“裸写”PyTorch 线性回归教学代码**,清晰地展示了深度学习训练的“前向 - 损失 - 反向 - 更新”四步流程。
&spm=1001.2101.3001.5002&articleId=158468125&d=1&t=3&u=b9768719804949348b661b7a6a510ff4)
2617

被折叠的 条评论
为什么被折叠?



