本节摘要:PyTorch 模型几乎都继承
nn.Module。__init__里创建nn.Linear、nn.ReLU等子模块并赋给self,forward里描述张量怎么流过它们。原文入门分类器:view成 784,Linear 到 128,ReLU,Linear 到 10,输出 logits。model.to(device)把参数迁到 GPU;model.parameters()交给 Adam;model.train()与model.eval()切换 Dropout 等行为。漏把层挂到self上,优化器就看不见那些权重。
阅读完本节,你应当能够:
nn.Module,并说明为何输出不加 Softmaxforward 与 __call__ 的关系:应调用 model(x) 而不是只写 forwardno_grad 的分工nn.Module 做了两件你不想手写的事:把赋给 self 的子模块和 Parameter 登记进树;提供 .to()、.parameters()、.state_dict()、.train() 这些树操作。forward 才是你真正要写的算法。原文强调:实现 __init__ 定义层,实现 forward 定义数据如何通过。不要在 forward 里用 nn.Linear(...) 当场新建层——那会每个批次创建新权重,学不会,显存还会涨。层必须在 __init__ 建一次。
调用模型写 outputs = model(inputs)。这会走 __call__,其中除了 forward 还有钩子。直接 model.forward(inputs) 会跳过钩子,评估和某些回调会 silently 失效。对照 Keras:model(x) 或 model.predict 都不要绕过公开入口。
原文 SimpleNN / SimpleClassifier 一类命名:self.fc1 = nn.Linear(28*28, 128),self.relu = nn.ReLU(),self.fc2 = nn.Linear(128, 10)。ReLU 无参数,写成模块或写成 F.relu 都行。有参数的层必须是模块属性。列表里的层要用 ModuleList,普通 Python list 不会登记参数——这是从 Keras Sequential 列表抄过来的人第二常见的坑(第一是 Softmax 加倍)。
💡 关键直觉:优化器只能更新它拿到的参数。
parameters()来自模块树,树来自self.xxx = nn.Something。没挂上树的张量就是装饰品。
原文把模型迁到 cuda:0 若可用否则 CPU,并打印 Using device。这行应出现在训练循环之前。之后每一批数据同样迁移。模型在 GPU、数据在 CPU 的错误信息直白,但有时只发生在某一层(部分子模块忘了 to),更难查。养成「构造完立刻 to,循环里数据立刻 to」的对称习惯。
model.train() 在每个训练 epoch 开头调用;model.eval() 在评估开头调用。对纯 Linear+ReLU 网络,两者前向数值相同。一旦有 Dropout 或 BN,不同。原文 4.2 示例包含模式切换,即使你当前网络没有 Dropout,也建议写上,形成模板。4.3 的循环会重复这套模板。
CrossEntropyLoss 吃 logits 与类别索引。forward 返回 fc2 的输出,不加 Softmax。推理若要概率,再 torch.softmax(logits, dim=1)。训练不要为了「好看」把 Softmax 写进 forward 还继续用交叉熵。Keras 入门把 Softmax 写进模型是因为损失默认吃概率;两边各自自洽。移植代码时改损失或改输出,只改一头。
| 习惯 | 做法 | 若忽略 |
|---|---|---|
| 层建在 init | self.fc = nn.Linear |
参数不更新或每步新建 |
| 调用模型 | model(x) |
钩子被跳过 |
| 输出 logits | 末层线性 | 与 CrossEntropy 双 Softmax |
| 设备对称 | 模型与批次同 to | 运行时报设备不匹配 |
| 模式 | train 对 eval | BN/Dropout 行为错 |
打印结构:print(model) 递归显示子模块。参数量:对 p.numel() 求和。与 Keras summary 对照,应几乎同一数字(偏置都算上)。差很多先查是否 Flatten 维写错,或是否多写了 Softmax 层参数(Softmax 无参数,差很多更可能是多了一层 Linear)。
函数式 nn.functional 与模块的分工:无状态用 F.relu 很常见;有状态必须模块。Dropout 用模块才能受 eval 控制。全用 Sequential 包起来也可以:self.net = nn.Sequential(nn.Flatten(), nn.Linear(784,128), nn.ReLU(), nn.Linear(128,10)),forward 一行。这和 Keras Sequential 最像。原文选择显式 fc1/fc2,便于在中间插 view。两种都对。
把层放进普通 list:self.layers = [nn.Linear(...), nn.Linear(...)],然后在 forward 里循环。看起来聪明,parameters() 是空的或不全。应 self.layers = nn.ModuleList([...])。字典用 ModuleDict。这是 Module 树的语法,没有对应 Keras 坑,因为 Keras Sequential 本来就只接受层对象列表且会登记。
共享权重:在 forward 里两次调用同一个 self.fc,才是共享。新建两个 Linear 即使形状相同也不共享。Keras 函数式里重复调用同一层对象同样是共享。互译时认「同一对象两次调用」。
⚠️ 常见坑:在
__init__里写self.weight = torch.zeros(...)却不用nn.Parameter包起来,权重不会出现在parameters()里,也不会进state_dict的常规路径。自定义层时用nn.Parameter。
forward 里的 Python if 就是动态图。按输入选择分支,梯度只沿实际路径走。Keras Sequential 没有这个位置给你写 if;函数式是静态接线;子类化 call 才能比。所以「复杂控制流」优势是真的,但只在你真的写了 if 时存在。Fashion MLP 没有 if,优势为零。不要用未使用的能力为选型辩护。
把 Keras Sequential 翻译成 Module 的机械步骤:每一层变成一个 self 属性;Flatten 变成 torch.flatten 或 view;带 activation 的 Dense 拆成 Linear+激活;末层 Softmax 删掉若损失是 CrossEntropy。反向翻译:把 Linear+ReLU 合成 Dense(activation='relu'),末层加 softmax,compile 稀疏交叉熵。第 5.2 节会按这个机械步骤给双侧代码。
设备与 state_dict 加载有关:保存时在 GPU 上的张量,加载到无 GPU 的机器要 map_location。4.4 展开。构建节记住:打印 next(model.parameters()).device 可知模型在哪,比猜强。
最后,模块可以嵌套:一个 Module 里 self.block = AnotherModule()。树是递归的。Keras 把自定义层继承 Layer。对照到写 Block 时再深入。入门一个类足够。
优化器可以接收参数组:骨干用较小学习率,分类头用较大。入门只有一组 model.parameters()。一旦你冻骨干,应把 filter(lambda p: p.requires_grad, model.parameters()) 交给 Adam,避免无梯度参数进入状态。Keras 冻层后 compile 一次更稳妥。对照到迁移学习时再启用,本课随机初始化不需要。
钩子(forward hook)能在不改 forward 的情况下偷看中间激活,适合查哪一层输出爆炸。Keras 也有类似回调。入门用打印 x.mean() 插在 forward 里更直接,查完删掉。不要把调试打印留进「最终模型」,否则 eval 时刷屏,还可能因打印触发同步让速度变差,误判框架慢。
nn.Sequential 嵌套与手写 fc1/fc2 在 state_dict 的 key 上不同:net.0.weight 对 fc1.weight。加载旧检查点时 key 对不上,不是权重坏了,是树的命名变了。改结构后不要强行 strict=False 吞错,除非你清楚哪些层是新的。4.4 会把这一点变成保存纪律。
Python 循环在层数固定且很少时完全可接受。真正慢的是对 batch 维写 Python for 逐张图算。向量化到批次维,让 Linear 与 Conv 吃整批。层数上的 for(例如重复同一个 block 四次)只要每次调用的是模块对象,仍然走高效核。不要因为「听说 Python 慢」把 128 单元 MLP 改成无法阅读的一长串。可读的 forward 比微不足道的解释器开销重要。
把 nn.Module 的注册规则当成一次「户口登记」。赋给 self 的模块会领到户口,出现在 parameters() 和 state_dict 里;塞进普通 list 的模块是黑户,前向或许还能跑(如果你手动调用了),优化器却看不见。Keras Sequential 接收的本来就是层对象列表,登记由容器完成,所以从 Keras 转过来的人最容易在这里翻车。对照检查:打印 list(model.named_parameters()) 的名字,应看到 fc1.weight、fc1.bias、fc2.weight、fc2.bias。少一对,说明那一层没挂上树。ReLU 模块无参数,名字可能不出现在 parameters 里,这是正常的,不要为了「对称」给 ReLU 加假权重。
设备迁移是树操作:model.to(device) 递归搬所有已登记参数。之后新建的张量不会自动跟着搬,所以循环里的批次仍要单独 to。有人在 __init__ 里缓存了 torch.zeros(1) 当缓冲区却忘了登记为 buffer,to 不会搬它,运行时设备和参数不一致。入门 MLP 不要缓存这种张量。需要缓存时用 register_buffer,它进 state_dict 但不进优化器,适合 BN 的滑动平均——又一条「等你加 BN 才需要」的伏笔。
原文把模型实例化后立刻判断 CUDA 是否可用,打印 Using device。请把这三行当成模板的固定抬头,不要凭感觉「我这台电脑没有显卡所以删掉」。删掉后,以后插上显卡你要改很多处。device 变量存在时,CPU 与 GPU 走同一套代码。这是和 Keras fit 内部处理设备相对应的显式版本:能力相同,切口不同。
Module 验收:named_parameters 里能看到 fc1 与 fc2 的 weight 和 bias;forward 返回未经过 Softmax 的十维;调用走 model(x) 而不是只调 forward;device 与即将到来的批次一致;train 与 eval 方法已经写进循环模板即使当前没有 Dropout。假前向 zeros(2,1,28,28) 得到 (2,10)。ModuleList 而不是 list。Parameter 而不是裸 zeros。这些检查比看准确率早一个数量级。和 Keras 互译时末层 Softmax 按损失删除或补上,不要两边都「为了好看」留着。打印模型结构字符串,确认没有意外多出来的线性层。
再核对一次 forward 与损失的边界。有人把 softmax 写进 forward 是为了 predict 方便,于是训练时改用 NLLLoss 去配已经取过对数的概率,结果和原文 CrossEntropyLoss 又错开一档。本课规定:forward 只出 logits,需要概率时在推理单元格临时 softmax。这样训练代码与原文一致,推理也不丢概率解释。Keras 入门网络把 softmax 放进模型,是因为它的损失默认吃概率;移植时改一头即可。named_parameters 的名字还决定检查点 key,改 fc1 为 linear1 会导致旧字典加载失败,这不是权重损坏,是户口本改名。改名后要么映射 key,要么当新实验重训。不要 strict=False 混过去还报告测试准确率。device 打印用 next(model.parameters()).device,比猜「我调用过 to 了」可靠。优化器必须在 to 之后创建,否则 Adam 缓冲可能留在 CPU,第一步更新就会设备不一致。这条顺序:先建模型,再 to,再建优化器,再进循环。写进模板抬头,4.3 会重复,这里先钉死。
把「优化器在 to 之后创建」写成代码审查的第一句。Adam 内部缓冲与参数同设备,先建优化器再 to 模型,缓冲可能留在旧设备。Keras compile 通常在模型已放好之后,较少踩这一下。PyTorch 用户从 Keras 转来最容易按「先把一切 new 出来再搬」的习惯踩雷。第二句:改完类定义必须 new 一个新实例,旧实例的层对象还是旧的。交互式笔记本尤其如此。第三句:print named_parameters 当作 summary。三句话构成 Module 的日常卫生,比任何技巧更能减少 5.2 的假对照。卫生没做就去调 Dropout,等于在脏盘子上摆盘。forward 里不要捕获外部全局张量当权重,那不会进 state_dict,保存时丢失,加载后那部分变成随机,准确率掉成谜。权重只从 self 上的 Parameter 来。
把户口登记、设备顺序、假前向三件事当成 Module 的晨检。晨检不过不去调参。晨检过了,5.2 只是把 zeros 换成真图。真图带来的新问题应只有数据契约,不应再有未登记的层。若 5.2 才发现 parameters 是空的,说明晨检被跳过。跳过晨检的对照数字,本课视为无效。
审查 named_parameters 时朗读每一个名字。读到不认识的名字,说明你复制了额外模块。读不到 fc2,说明输出层没挂上。朗读比看长度有效,因为长度对了名字仍可能错。错名字会在加载旧检查点时以 key 错误出现,那时已经晚了。晨检把错误提前到声明当天。当天修,比第 5 章修便宜一个数量级。便宜的原因是此时还没有曲线可以让你舍不得推倒重来。
Module 节结束前再假前向一次,输出必须是批次乘十。不是十,就是类别维丢了。丢了还能算损失,但那是错的损失。错的损失会在 5.2 给你一条下降的曲线,下降不能证明类别维还在。假前向证明它还在。证明了再毕业。毕业证是假前向那一行输出形状,不是准确率。准确率是第 5 章的事,形状是本节的事。两件事不要抢同一行日志。
下一节把 Keras 三板斧和 PyTorch 五步循环并排,完成训练对照轴。