沐神-动手学习深度学习,课后习题 3.3 线性回归的简洁实现
练习问题:
1.如果将小批量的总损失替换为小批量损失的平均值,需要如何更改学习率?
2.查看深度学习框架文档,它们提供了哪些损失函数和初始化方法?用Huber损失代替原损失,即 (3.3.1) 其 它 情 况 l ( y , y ′ ) = { | y − y ′ | − σ 2 if | y − y ′ | > σ 1 2 σ ( y − y ′ ) 2 其它情况
3.如何访问线性回归的梯度?
答案
根据上传的文档内容,以下是针对“3.3. 线性回归的简洁实现”章节末尾练习问题的解答:
1. 如果将小批量的总损失替换为小批量损失的平均值,需要如何更改学习率?
在原始代码中,损失函数 nn.MSELoss() 默认返回所有样本损失的平均值(即均方误差的平均值)CITE12CITE_{12}CITE12。
若改为使用小批量的总损失(即损失和),需调整学习率:
- 原始设定:损失为平均值,学习率 lr=0.03lr=0.03lr=0.03。
- 修改后:若损失改为总和 Lsum=∑i=1nliL_{\text{sum}} = \sum_{i=1}^n l_iLsum=∑i=1nli,梯度计算会扩大 nnn 倍(nnn 为批次大小),因为 ∇Lsum=n⋅∇Lavg\nabla L_{\text{sum}} = n \cdot \nabla L_{\text{avg}}∇Lsum=n⋅∇Lavg。
- 学习率调整:需将学习率缩小 nnn 倍,即 lrnew=lrnlr_{\text{new}} = \frac{lr}{n}lrnew=nlr。例如当批次大小 n=10n=10n=10 时,新学习率应为 0.0030.0030.003 CITE13CITE_{13}CITE13。
2. 用Huber损失替代原损失函数
(1) 深度学习框架提供的损失函数和初始化方法
- 损失函数:PyTorch 提供
MSELoss(均方误差)、L1Loss(绝对值损失)、HuberLoss、CrossEntropyLoss等 CITE12CITE_{12}CITE12。 - 初始化方法:包括
normal_(正态分布初始化)、fill_(常量初始化)、uniform_(均匀分布初始化)等 CITE11CITE_{11}CITE11。
(2) Huber损失公式
Huber损失结合了L1和L2损失的优点,对异常值更鲁棒:
[
l(y, y’) =
\begin{cases}
|y - y’| - \frac{\sigma}{2} & \text{if } |y - y’| > \sigma \
\frac{1}{2\sigma} (y - y’)^2 & \text{其他情况}
\end{cases}
]
其中 σ\sigmaσ 是阈值参数 CITE18CITE_{18}CITE18。
(3) 代码实现替换
# 原损失函数(均方误差)
loss = nn.MSELoss()
# 替换为Huber损失
loss = nn.HuberLoss(delta=sigma) # delta 即公式中的 σ
训练时直接使用 nn.HuberLoss() 替代 nn.MSELoss() 即可 CITE12CITE_{12}CITE12。
3. 如何访问线性回归的梯度?
在PyTorch中,梯度通过自动微分机制存储,可通过以下步骤访问:
- 前向传播:计算损失
l = loss(net(X), y)。 - 反向传播:调用
l.backward()计算梯度。 - 访问梯度:
# 获取权重梯度 weight_grad = net[0].weight.grad # 获取偏置梯度 bias_grad = net[0].bias.grad
梯度存储在张量的 .grad 属性中,需在调用 backward() 后访问 CITE14CITE_{14}CITE14。
注意:每次更新前需调用
trainer.zero_grad()清除历史梯度,避免累加 CITE14CITE_{14}CITE14。
关键引用说明
- CITE11CITE_{11}CITE11:参数初始化方法(
normal_和fill_)CITE11CITE_{11}CITE11。 - CITE12CITE_{12}CITE12:损失函数定义(
MSELoss和HuberLoss)CITE12CITE_{12}CITE12。 - CITE13CITE_{13}CITE13:优化器学习率与批次大小的关系CITE13CITE_{13}CITE13。
- CITE14CITE_{14}CITE14:梯度访问与训练流程CITE14CITE_{14}CITE14。
- CITE18CITE_{18}CITE18:Huber损失公式CITE18CITE_{18}CITE18。
更多推荐

所有评论(0)