2.8 Chapter 2 练习题: PyTorch 基础与训练流程

Author

xcherry9988-code

Published

2026-09-21

Modified

2026-09-21

练习题

1. 对于最常见的 map-style Dataset,通常需要实现 ______ 和 ______ 两个方法;而 DataLoader 的主要作用是把多个样本组织成可以用于训练的 ______。

2. 关于 model.eval() 和 torch.no_grad(),下列说法正确的是:

A. model.eval() 会自动关闭 Autograd,因此验证时不再需要 torch.no_grad()。
B. torch.no_grad() 会把模型中的 Dropout 和 BatchNorm 自动切换到评估模式。
C. model.eval() 控制某些模块在训练和评估阶段的行为,而 torch.no_grad() 控制当前这段计算是否被 Autograd 记录。
D. 两者作用完全相同,只是写法不同。

3. 一个最基本的 PyTorch 参数更新过程通常包含:

optimizer.______()
loss.______()
optimizer.______()

其中,第一步用于清理旧梯度,第二步计算当前 loss 对参数的梯度,第三步真正修改参数。

4. 使用 Adam 优化器训练模型时,如果希望程序中断后尽可能从原来的训练状态继续,下面哪组内容最应该保存在 checkpoint 中?

A. 只保存 model.state_dict()。
B. 保存 model.state_dict()、optimizer.state_dict() 和当前训练进度,例如 epoch。
C. 只保存 optimizer.state_dict()。
D. 只保存训练集和每个 epoch 的 loss。

5. 运行下面的代码:

import torch
import torch.nn as nn


class Demo(nn.Module):
    def __init__(self):
        super().__init__()
        self.tensor = torch.randn(2, 2, requires_grad=True)
        self.param = nn.Parameter(torch.randn(2, 2))


model = Demo()
print([name for name, _ in model.named_parameters()])

写出输出,并解释为什么 tensor 和 param 都可以设置梯度属性,但只有其中一个会被 parameters() 和 named_parameters() 找到。

参考答案

1. 答案:len();getitem();mini-batch

Dataset 负责定义数据集中有多少样本,以及如何根据下标取出一个样本;DataLoader 再负责取出多个样本并组织成 mini-batch。

2. 答案:C

model.eval() 控制的是模块行为,例如关闭 Dropout 的随机丢弃,并让 BatchNorm 使用评估阶段的统计方式,但它不会关闭自动微分;torch.no_grad() 控制的是 Autograd 是否记录当前计算。

验证阶段通常需要同时使用:

model.eval()

with torch.no_grad():
    ...

3. 答案:zero_grad;backward;step

完整形式为:

optimizer.zero_grad()
loss.backward()
optimizer.step()

PyTorch 中梯度默认会累积。因此通常先清理旧梯度,然后使用 backward() 计算梯度,最后的 step() 才真正根据梯度和优化器状态更新参数。

4. 答案:B

如果只是推理,只保存模型参数通常已经足够。但如果目标是中断后继续训练,还应该恢复优化器内部状态以及当前训练进度,例如 Adam 会保存一阶矩、二阶矩等历史状态,只恢复模型参数会改变后续优化轨迹。

5. 答案:[‘param’]

self.tensor 虽然设置了 requires_grad=True,可以参与 Autograd 并计算梯度,但它仍然只是一个普通 Tensor,不会因为被赋值给 Module 属性就自动注册成模型参数。而 self.param 是 Parameter。当 Parameter 被赋值给 Module 的属性时,PyTorch 会自动将它注册为模型参数,因此能够被 model.parameters() 和 model.named_parameters() 找到,也可以自动交给优化器管理。