2.2 用 nn.Module 组装远征队


文档摘要

2.2 用 nn.Module 组装远征队 本节摘要:nn.Module 是 PyTorch 模型的唯一组织形式。本节讲透三段式写法(声明子层、描述数据流、隐式调用),用一次"参数失踪案"揭示参数登记机制,并组装出贯穿全册的手写数字识别模型。 组队规则:三段式 上一节你手写了矩阵运算,也见过了 nn.Linear。本节把视角抬高一层:怎么把一堆层组织成一个模型。PyTorch 的答案是所有模型继承 ,写法固定为三段—— 里把用到的子层声明成属性, 里描述数据怎么流过这些子层,外部直接 拿到输出(不必写 ,因为基类在背后做了挂钩)。 三段式不是风格偏好,而是框架的登记机制:凡是 里赋值给 的 与 ,都会被自动登记进模型的成员表。

2.2 用 nn.Module 组装远征队

本节摘要:nn.Module 是 PyTorch 模型的唯一组织形式。本节讲透三段式写法(声明子层、描述数据流、隐式调用),用一次"参数失踪案"揭示参数登记机制,并组装出贯穿全册的手写数字识别模型。

组队规则:三段式

上一节你手写了矩阵运算,也见过了 nn.Linear。本节把视角抬高一层:怎么把一堆层组织成一个模型。PyTorch 的答案是所有模型继承 nn.Module,写法固定为三段——__init__ 里把用到的子层声明成属性,forward 里描述数据怎么流过这些子层,外部直接 model(x) 拿到输出(不必写 model.forward(x),因为基类在背后做了挂钩)。

三段式不是风格偏好,而是框架的登记机制:凡是 __init__ 里赋值给 selfnn.Modulenn.Parameter,都会被自动登记进模型的成员表。优化器靠这张表找参数,保存加载靠这张表存状态,model.train()model.eval() 也靠它把模式广播给所有子层。理解了"登记",第 5 章的很多现象就顺理成章了。

图 2-1:nn.Module 的登记机制

图 2-1:nn.Module 的登记机制

组装贯穿案例的模型

动手写贯穿全册的那个模型。尺寸刻意小:输入 64 维(8×8 图像展平),一个隐藏层 32 个单元,输出 10 类:

import torch import torch.nn as nn class DigitNet(nn.Module): """8x8 手写数字识别的多层感知机,本册贯穿案例的主角""" def __init__(self, in_dim=64, hidden=32, n_class=10): super().__init__() # 先完成基类初始化,漏写会报错 self.flatten = nn.Flatten() # 64x1x8x8 -> 64x64 self.fc1 = nn.Linear(in_dim, hidden) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden, n_class) def forward(self, x): x = self.flatten(x) x = self.relu(self.fc1(x)) return self.fc2(x) # 输出 logits,不做 softmax model = DigitNet() print(model) print("可训练参数总数:", sum(p.numel() for p in model.parameters()))

输出:

DigitNet( (flatten): Flatten(start_dim=1, end_dim=-1) (fc1): Linear(in_features=64, out_features=32, bias=True) (relu): ReLU() (fc2): Linear(in_features=32, out_features=10, bias=True) ) 可训练参数总数: 2442

2442 与 2.1 节的手算完全对上(64×32+32 加 32×10+10)。两个写法要点:super().__init__() 必须第一行调用,否则登记表没建好,后续赋值会崩;forward 末尾返回的是 logits 而非概率,softmax 留给损失函数内部处理——第 3 章讲交叉熵时解释为什么这样分工。

完整案例:参数失踪案

背景:某同事想让模型学一个"温度系数"来缩放输出,在 __init__ 里写了 self.temperature = torch.ones(1),训练完发现这个值纹丝不动,优化器好像从没见过它。

操作:对比两种写法在参数表里的存在性。

class Broken(nn.Module): def __init__(self): super().__init__() self.temperature = torch.ones(1) # 普通张量:不入册 class Fixed(nn.Module): def __init__(self): super().__init__() self.temperature = nn.Parameter(torch.ones(1)) # nn.Parameter:入册 b, f = Broken(), Fixed() print("Broken 参数:", list(b.named_parameters())) print("Broken 里温度可求导:", b.temperature.requires_grad) print("Fixed 参数:", [(n, tuple(p.shape)) for n, p in f.named_parameters()]) print("Fixed 里温度可求导:", f.temperature.requires_grad)

输出:

Broken 参数: [] Broken 里温度可求导: True Fixed 参数: [('temperature', (1,))] Fixed 里温度可求导: True

结果:两种写法下张量都设了可求导,但只有 nn.Parameter 版本出现在参数表里。

解读:requires_grad=True 只保证"梯度会被算出来",而优化器只认 model.parameters() 这张名册。普通张量即使有梯度,优化器也无从知晓,于是永远不被更新。这就是"登记机制"的边界:入册靠类型,不靠求导标记。

变式:如果这个温度参数你希望固定不学(比如推理时手动调),正确做法反过来——用普通张量并配 register_buffer,这样它不出现在参数表里、却会随 state_dict 一起保存。这个三态区分(普通张量、Parameter、buffer)是读成熟模型源码的必修语法。

本节要点回顾

  • 三段式__init__ 声明子层、forward 描述数据流、model(x) 隐式调用;
  • 登记机制是 nn.Module 的灵魂:赋值即入册,优化器、保存、模式切换全靠这张表;
  • 参数失踪案的病根:普通张量不入册,优化器查无此人,须用 nn.Parameter
  • forward 输出 logits,softmax 归损失函数管。

队伍集结完毕。下一节带你在预置层的货架上巡一圈,认识那些会在实战里高频出面的标准成员。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U