Skip to content

09. PyTorch nn Module Basics | PyTorch nn.Module 基础 ​

难度: Medium | 环境: CPU-first | 标签: PyTorch, 模型封装, 状态管理 | 目标人群: Part 2-4 前置补课者

🚀 云端运行环境

本章节的实战代码可以点击以下链接在免费 GPU 算力平台上直接运行:

Open In ColabOpen In Studio (国内推荐:魔搭社区免费实例)

本页聚焦:会用 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 的持久化细节,再顺着看下面这一页。

Q1:如何把模型组织成一个可调用的 nn.Module? ​

模型不是一堆散落函数,而是一个有边界、可调用、可递归组合的对象。nn.Module 先提供三个入口:用 __init__() 注册参数和子模块,用 forward() 描述计算,用 model(x) 触发前向。状态保存放到 Q4 再处理。

python
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 和参数总量是否符合预期。

python
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、统计量、固定缩放值
普通属性否否否配置、名称、临时标记
python
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,但不会把普通属性当成可保存状态。

python
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() 对照模块层级。这里关注的是数据流和模块边界,不是记住每个接口的名称。

python
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 和模块结构,确认组合后的网络还能正常执行。

python
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、学习率和训练步数。

python
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 的值还在不在。你要把它理解成最小的“状态闭环”:只要状态没变,模型就能被稳定恢复。

python
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 保持一致')

Released under the MIT License.