练习问题:

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=nLavg
  • 学习率调整:需将学习率缩小 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(绝对值损失)、HuberLossCrossEntropyLossCITE12CITE_{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中,梯度通过自动微分机制存储,可通过以下步骤访问:

  1. 前向传播:计算损失 l = loss(net(X), y)
  2. 反向传播:调用 l.backward() 计算梯度。
  3. 访问梯度
    # 获取权重梯度
    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:损失函数定义(MSELossHuberLossCITE12CITE_{12}CITE12
  • CITE13CITE_{13}CITE13:优化器学习率与批次大小的关系CITE13CITE_{13}CITE13
  • CITE14CITE_{14}CITE14:梯度访问与训练流程CITE14CITE_{14}CITE14
  • CITE18CITE_{18}CITE18:Huber损失公式CITE18CITE_{18}CITE18
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐