09. PyTorch nn Module Basics | PyTorch nn.Module 基础
难度: Medium | 环境: CPU-first | 标签: PyTorch, 模型封装, 状态管理 | 目标人群: Part 2-4 前置补课者
🚀 云端运行环境
本章节的实战代码可以点击以下链接在免费 GPU 算力平台上直接运行:
本页聚焦:会用 nn.Module 封装最小模型;会区分参数、缓冲区和普通属性;会读懂 state_dict() 和 parameters()。这一页开始把前面的 Tensor 和 Autograd 结果,正式收进一个模型对象里,让“能算”变成“能组织成模块”。
阅读顺序可以按这条线走:先看模型对象怎么成立,再看状态边界怎么划分,然后看 forward() 怎么串子模块,最后看 state_dict() 怎么保存和恢复。
关键词: nn.Module, Parameter, state_dict
显存路线视角: 参数、buffer 和普通属性的注册方式决定哪些对象会进入模型状态、迁移到 device 或被保存;本页只验证模块和状态边界,不估算完整训练峰值。完整账本还要把参数、梯度和 optimizer state 放在一起,可参阅 20 显存账本。
前置阅读
导语: 先看 0C 组页,把梯度边界和模型封装的边界对齐,再进入这一页会更顺。
相关阅读
导语: 本页先把 nn.Module、参数注册和状态保存的最小判断讲清楚;如果想继续看 state_dict 的持久化细节,再顺着看下面这一页。
- 10. PyTorch State_dict and Persistence | PyTorch 状态管理与持久化
- 20. Profiling and Memory Ledger | 性能剖析与显存账本
Q1:如何把模型组织成一个可调用的 nn.Module?
模型不是一堆散落函数,而是一个有边界、可调用、可递归组合的对象。nn.Module 先提供三个入口:用 __init__() 注册参数和子模块,用 forward() 描述计算,用 model(x) 触发前向。状态保存放到 Q4 再处理。
import torch
import torch.nn as nn
class SimpleLinear(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
# `Parameter` 会被自动注册进模型参数里。
self.weight = nn.Parameter(torch.empty(out_features, in_features))
self.bias = nn.Parameter(torch.zeros(out_features))
nn.init.xavier_uniform_(self.weight)
def forward(self, x):
# `forward()` 定义计算逻辑,真正调用时写 `model(x)`。
return x @ self.weight.t() + self.bias
class TwoLayerMLP(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super().__init__()
# 子模块会递归参与参数注册和状态管理。
self.fc1 = SimpleLinear(input_dim, hidden_dim)
self.act = nn.GELU()
self.fc2 = SimpleLinear(hidden_dim, output_dim)
def forward(self, x):
# 典型的模块组合:线性层 -> 激活 -> 线性层。
return self.fc2(self.act(self.fc1(x)))
def count_parameters(module):
"""统计模块及其子模块中已注册参数的总元素数。"""
return sum(param.numel() for param in module.parameters())
model = TwoLayerMLP(4, 8, 2)
print('参数总量:', count_parameters(model))
print('named_parameters:', [name for name, _ in model.named_parameters()])Q1验证:参数注册和前向输出是否正确?
这里直接检查两件事:自定义层能不能跑,MLP 的输出 shape 和参数总量是否符合预期。
def test_simple_linear():
layer = SimpleLinear(2, 1)
with torch.no_grad():
layer.weight.copy_(torch.tensor([[2.0, -1.0]]))
layer.bias.copy_(torch.tensor([0.5]))
x = torch.tensor([[3.0, 4.0]])
y = layer(x)
assert torch.allclose(y, torch.tensor([[2.5]]))
print('✅ SimpleLinear 通过')
def test_two_layer_mlp():
model = TwoLayerMLP(4, 8, 2)
x = torch.randn(3, 4)
y = model(x)
assert y.shape == (3, 2)
expected_params = (4 * 8 + 8) + (8 * 2 + 2)
assert count_parameters(model) == expected_params
print('✅ TwoLayerMLP 通过')
test_simple_linear()
test_two_layer_mlp()Q2:Parameter、buffer 和普通属性分别会进入哪里?
阅读一个模块时,先用“是否训练、是否保存、是否随模块迁移”三件事判断对象的归属。下面的表格给出本节需要掌握的最小区别。
| 对象 | 会被优化器更新 | 进入 state_dict() | 随模块迁移 device | 常见用途 |
|---|---|---|---|---|
Parameter | 是 | 是 | 是 | 权重、偏置 |
buffer | 否 | 是 | 是 | mask、统计量、固定缩放值 |
| 普通属性 | 否 | 否 | 否 | 配置、名称、临时标记 |
class ToyModule(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.tensor([1.0]))
# `buffer` 常用于 mask、统计量、缓存这类“要保存但不训练”的状态。
self.register_buffer('scale', torch.tensor([2.0]))
# 普通属性不会进入 state_dict,也不会随模块自动迁移到其他 device。
self.name = 'toy'
def forward(self, x):
return x * self.weight * self.scale
m = ToyModule()
print('parameters:', [n for n, _ in m.named_parameters()])
print('buffers:', [n for n, _ in m.named_buffers()])
print('state_dict keys:', list(m.state_dict().keys()))Q2验证:state_dict 是否同时包含参数和 buffer?
这里确认 state_dict() 会包含参数和 buffer,但不会把普通属性当成可保存状态。
m = ToyModule()
keys = list(m.state_dict().keys())
assert 'weight' in keys
assert 'scale' in keys
assert 'name' not in keys
print('✅ state_dict 结构通过')Q3:如何阅读 forward() 和子模块组合?
先沿着 forward() 看数据经过哪些子模块,再用 named_children() 或 named_modules() 对照模块层级。这里关注的是数据流和模块边界,不是记住每个接口的名称。
x = torch.randn(2, 4)
model = TwoLayerMLP(4, 8, 2)
# 调用模块时写 `model(x)`,PyTorch 会帮你走到 `forward()`。
y = model(x)
print('输出 shape:', y.shape)
print('模块层级:', model)
print('named_children:', [name for name, _ in model.named_children()])Q3验证:子模块组合是否正常?
这里直接检查前向输出 shape 和模块结构,确认组合后的网络还能正常执行。
model = TwoLayerMLP(4, 8, 2)
x = torch.randn(5, 4)
y = model(x)
assert y.shape == (5, 2)
print('✅ forward 和子模块组合通过')Q4:如何完成模型状态的最小保存与恢复?
模型结构确定后,再把参数和 buffer 做一次最小保存—恢复闭环。state_dict() 负责取得当前模型状态,load_state_dict() 负责恢复这些状态;它们处理的是模型状态,不包含 optimizer、学习率和训练步数。
model = TwoLayerMLP(3, 6, 2)
# `state_dict()` 只保存状态,不保存训练过程。
before = {k: v.clone() for k, v in model.state_dict().items()}
buffer = model.state_dict()
model.load_state_dict(buffer)
after = model.state_dict()
for key in before:
assert torch.allclose(before[key], after[key])
print('✅ state_dict roundtrip 通过')Q4验证:保存和恢复是否不丢状态?
这里只检查一件事:读出来再写回去,参数和 buffer 的值还在不在。你要把它理解成最小的“状态闭环”:只要状态没变,模型就能被稳定恢复。
model = TwoLayerMLP(3, 6, 2)
before = {k: v.clone() for k, v in model.state_dict().items()}
model.load_state_dict(model.state_dict())
after = model.state_dict()
for key in before:
assert torch.allclose(before[key], after[key])
print('✅ state_dict 保持一致')