1. 先搞清楚torch.nn.Module这个类到底在管什么凡是写过PyTorch代码的人第一行import大概率就是import torch.nn as nn。但说句实在话很多人用nn.Module只是照葫芦画瓢——class里写个__init__super一下再写个forward就完事了。这个类为什么叫Module、底层帮你做了什么、为什么所有网络结构都要继承它可能就讲不太清楚了。其实你可以把nn.Module理解成一个万能容器管家。它本身不存任何具体的网络逻辑但它管着你要搭的每一层网络背后那堆繁琐的底层事务参数怎么注册、梯度怎么传播、模块怎么嵌套、模型怎么保存加载、CPU和GPU之间怎么搬家。换句话说nn.Module是PyTorch深度学习代码的地基你所有自定义的网络层、完整的模型、甚至一个复杂的训练流程里那些可学习的状态最终都得靠它来承载。这一篇打算把nn.Module从里到外拆一遍包括它的设计逻辑、几个不能不知道的内部机制、那些藏在源码角落里却天天影响你的细节以及我实际写代码时踩过的一些坑。不管你是刚看完教程准备写第一个CNN的新手还是已经写了半年模型但总觉得继承一下就行的进阶玩家这篇都应该有值得你收藏的内容。2. nn.Module的核心机制——它究竟替你做了什么2.1 参数注册为什么变量要放进self.xxx nn.Parameter(...)而不是直接赋值先看一个最简单但很典型的例子import torch import torch.nn as nn class MyLinear(nn.Module): def __init__(self, in_features, out_features): super().__init__() self.weight nn.Parameter(torch.randn(in_features, out_features)) self.bias nn.Parameter(torch.randn(out_features)) def forward(self, x): return x self.weight self.bias这里的关键不是forward里的矩阵乘法而是nn.Parameter这层包装。很多初学者会问我不就是用self.weight torch.randn(...)存下权重吗为什么非要包一层nn.Parameter答案在于nn.Module内部实现了一个特殊的__setattr__钩子。当你在Module实例上做属性赋值时它不会老老实实只存一个属性而是会检查赋进来的值是什么类型如果是nn.Parameter就把它登记进模块的_parameters这个有序字典里如果是nn.Module就登记进_modules如果是普通Tensor或其他对象才走正常的Python属性赋值。这个设计带来的直接好处就是你在训练代码里写model.parameters()就能拿到所有子模块、所有层里的全部可学习参数。如果不用nn.Parameter而是直接self.weight torch.randn(...)那么这个Tensor就只是Python对象上一个普通属性optimizer torch.optim.SGD(model.parameters(), lr...)的时候根本扫描不到它梯度更不会计算训练时这个权重就是一个死值。同样地nn.Module还提供了register_buffer()方法来注册非可学习的持久化状态比如BatchNorm里的running_mean和running_var。这些buffer不需要梯度但需要跟模型一起保存和迁移设备。class MyBatchNorm(nn.Module): def __init__(self, num_features): super().__init__() self.register_buffer(running_mean, torch.zeros(num_features)) self.register_buffer(running_var, torch.ones(num_features)) def forward(self, x): # 实际BN逻辑省略这里只是为了展示buffer return (x - self.running_mean) / torch.sqrt(self.running_var 1e-5)2.2forward与__call__框架把你调用的入口暗箱操作了nn.Module里另一个关键设计是你定义的是forward但调用的时候用的是model(x)。为什么不是直接调用model.forward(x)因为nn.Module的__call__方法内部做了一层包装。它会在真正执行forward前触发注册在模块上的前置钩子forward_pre_hooks在forward执行完之后再触发后置钩子forward_hooks。这些钩子机制是后面要做特征提取、梯度修改、模型诊断的基础。所以如果你写代码时绕开model(x)直接调model.forward(x)钩子不会生效模式管理model.train()/model.eval()相关的内部状态虽然还在但部分依赖钩子的功能会失效比如某些第三方库的hook逻辑。我见过有人为了性能优化直接调forward结果模型输出不对查了半天是忘了走__call__的钩子。注意自己定义子类的时候永远不要重写__call__。要改就改forward。PyTorch官方的设计就是让__call__承载通用逻辑forward承载模型特定逻辑。一旦重写__call__等于丢掉了Module自带的一整套钩子与管理机制。2.3state_dict与load_state_dict模型参数的快照机制训练完模型最标准的保存方式不是直接torch.save(model)而是保存model.state_dict()。这个state_dict本质上是一个OrderedDict里面包含所有参数和buffer的键值对。键名就是你从Module根节点出发的路径比如features.0.weight、classifier.3.bias。这种路径式的命名不是随便设计的它为后面做模型剪枝、迁移学习、模块替换提供了极大的便利。# 保存与加载模型只存参数不存结构 torch.save(model.state_dict(), model.pt) model.load_state_dict(torch.load(model.pt))我强烈建议你养成只保存state_dict的习惯而不是整个model对象。原因有两个第一直接torch.save(model)会把这个对象的所有Python属性都序列化进去文件体积更大而且一旦代码结构改了比如类名变了、层顺序调整了老模型文件很可能加载不回来。第二只保存参数的话你在加载时可以灵活地只加载一部分权重——这是做迁移学习的核心操作。比如你在大模型上训好了backbone现在要换一个分类头你可以构造新模型然后只加载匹配的层pretrained torch.load(pretrained.pt) model NewModel(num_classes10) # 过滤掉分类头相关的键 filtered {k: v for k, v in pretrained.items() if k.startswith(backbone.)} model.load_state_dict(filtered, strictFalse)strictFalse这个参数也很常用。它允许加载时忽略不匹配的键比如你自己改了某层输出维度导致fc.weight形状对不上strict模式下直接报错而strictFalse会跳过不匹配的部分并在返回值里告诉你哪些键缺失、哪些键不匹配。3. 嵌套的Module结构——Sequential、ModuleList、ModuleDict该怎么选3.1 三种容器类各自的使用场景一个完整的模型必然包含很多层如果全部手动在__init__里写成一长串属性代码会非常难看。nn.Module也提供了几个内置容器但很多人会把它们混用导致踩坑。先看它们是什么import torch.nn as nn # Sequential按顺序执行层之间是线性关系 model nn.Sequential( nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 10), ) # ModuleList像Python list一样存多个模块但不自动执行 layers nn.ModuleList([nn.Linear(64, 64) for _ in range(3)]) # ModuleDict像dict一样存模块用字符串做键 blocks nn.ModuleDict({ encoder: nn.Linear(64, 128), decoder: nn.Linear(128, 64), })关键的区别在于Sequential把模块按顺序串起来输入自动从第一个层流到最后一个层适合那种一条路走到底的网络结构比如基础的MLP、简单的CNN。而ModuleList和ModuleDict只是帮你收纳模块让模块能被model.parameters()识别但不负责前向传播——你得自己在forward里写遍历或按索引取用。举个例子同样是三层线性层用Sequential写就是self.layers nn.Sequential( nn.Linear(64, 64), nn.ReLU(), nn.Linear(64, 64), )而在forward里直接写x self.layers(x)就行了。但如果你的网络里有一个循环执行同一个block多次再叠加残差的结构用ModuleList更合适class ResBlock(nn.Module): def __init__(self, dim): super().__init__() self.net nn.Sequential( nn.Linear(dim, dim), nn.ReLU(), nn.Linear(dim, dim), ) def forward(self, x): return x self.net(x) class MyModel(nn.Module): def __init__(self, dim, num_blocks): super().__init__() self.blocks nn.ModuleList([ResBlock(dim) for _ in range(num_blocks)]) def forward(self, x): for block in self.blocks: x block(x) return x这里用ModuleList而不是Python原生list的关键原因在于如果你写self.blocks [ResBlock(dim) for _ in range(3)]这三个ResBlock不会被nn.Module扫描到它们的参数就不会出现在model.parameters()中反向传播时这些层的权重永远不更新——这是新手最容易犯的隐蔽错误之一。那么ModuleDict什么时候用典型场景是动态选择路径。比如你有一个模型根据输入类型选择不同的编码器class MultiEncoderModel(nn.Module): def __init__(self): super().__init__() self.encoders nn.ModuleDict({ image: nn.Linear(784, 256), text: nn.Linear(300, 256), }) def forward(self, x, modality): return self.encoders[modality](x)3.2 不推荐的做法与替换方案还有一个很多老代码里会出现的东西nn.Sequential里放OrderedDict。这个用法本身没问题但如果你对Python比较熟可能会想用普通dict来存模块。我劝你千万别这么做因为nn.Module不会自动把普通dict里的子模块注册进去这会导致一个非常难检查的问题——代码不报错但loss不下降。排查了半天最后发现是某个模块因为放在普通dict里参数压根没有参与优化。类似的坑还有直接在__init__里用setattr动态添加模块。虽然setattr也能触发nn.Module的注册机制但代码可读性会变得很差。如果确实需要动态添加子模块我更推荐先把结构设计清楚或者用ModuleDict这类容器统一管理。4. 手把手实现一个可复用的自定义Module4.1 自定义模块的完整代码前面把原理讲了不少现在来一个真正能直接用的例子。这个例子不是那种一眼就会的玩具Linear而是一个带有残差连接、可配置dropout、并且支持特征提取功能的MLP Block。这个结构在很多实际项目里都能直接抄。import torch import torch.nn as nn import torch.nn.functional as F class MLPBlock(nn.Module): def __init__(self, in_dim, hidden_dim, out_dim, dropout0.1): super().__init__() self.fc1 nn.Linear(in_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, out_dim) self.dropout nn.Dropout(dropout) self.shortcut None # 如果输入输出维度不一致需要加一个投影层做残差连接 if in_dim ! out_dim: self.shortcut nn.Linear(in_dim, out_dim) # 初始化权重这个别忘了 self._init_weights() def _init_weights(self): nn.init.xavier_uniform_(self.fc1.weight) nn.init.xavier_uniform_(self.fc2.weight) nn.init.zeros_(self.fc1.bias) nn.init.zeros_(self.fc2.bias) if self.shortcut is not None: nn.init.xavier_uniform_(self.shortcut.weight) nn.init.zeros_(self.shortcut.bias) def forward(self, x): identity x x self.fc1(x) x F.relu(x) x self.dropout(x) x self.fc2(x) if self.shortcut is not None: identity self.shortcut(identity) return x identity def forward_feature(self, x): # 返回倒数第二层的特征用于后续做特征可视化或对比学习 x self.fc1(x) x F.relu(x) x self.dropout(x) return x这个模块看起来不复杂但里面有几个细节是值得展开讲的。第一个细节是self.shortcut这个分支。它不是一开始就定义的而是根据in_dim和out_dim是否相等来决定的。如果相等我们可以直接做恒等映射如果不相等残差连接就得加一个线性投影把维度对齐。这种条件初始化的写法在Transformer类模型里特别常见比如BERT的不同层维度都一样所以不需要投影但当你把Encoder输出接到不同尺寸的Decoder时这个逻辑就会用到。第二个细节是独立的_init_weights方法。很多人写nn.Module子类时根本不初始化权重全靠PyTorch默认的初始化。说实话对于大部分场景默认初始化是能跑的但在深网络里不显式做Xavier或Kaiming初始化很可能会出现训练初期loss下降极慢甚至梯度爆炸的问题。自己写一个_init_weights放进__init__里是成本最低的稳定性保障。第三个细节是forward_feature这个方法。它不在forward里面执行完整的前向而是截断了最后一层返回隐藏层特征。这种做法的价值在于当你要做迁移学习、特征可视化、或者像SimCLR那样把中间层的表征拉出来做对比损失时不需要另外hack模型内部结构直接调用这个方法就行。4.2 使用这个自定义模块import torch.optim as optim model MLPBlock(in_dim128, hidden_dim256, out_dim64, dropout0.2) # 1. 查看模型结构和参数 print(model) print(sum(p.numel() for p in model.parameters())) # 输出30912128*256 256 256*64 64 128*64 64如果有shortcut还要加上投影层参数 # 2. 前向传播 x torch.randn(4, 128) y model(x) print(y.shape) # torch.Size([4, 64]) # 3. 提取中间特征 feat model.forward_feature(x) print(feat.shape) # torch.Size([4, 256]) # 4. 用Adam优化参数能被正常扫描到 optimizer optim.Adam(model.parameters(), lr1e-3)这里的参数量计算也是一个值得养成的习惯。128*256是fc1的weight256是fc1的bias256*64是fc2的weight64是fc2的biasshortcut投影层128*64加64。加起来是32768 256 16384 64 8192 64 57728。如果你发现自己实现的模型参数量跟理论值对不上多半就是有模块没被注册或者维度写错了。4.3train和eval模式到底切换了什么nn.Module还内置了train()和eval()两个方法。它们有什么用关键在于某些层在两种模式下行为不一样。nn.Dropout训练时随机丢弃一部分神经元eval时什么都不做nn.BatchNorm训练时用当前batch的均值和方差做归一化同时更新running_mean和running_vareval时直接用running统计量nn.Dropout2d、nn.AlphaDropout等类似。所以如果你有Dropout或BN层在推理前一定要记得调用model.eval()否则同样的输入每次输出的结果都会不一样Dropout的随机性或者推理结果跟用训练好的参数计算不一致BN用了错误的统计方式。相反训练之前要记得调回model.train()。model.eval() with torch.no_grad(): y model(x)这里还有个配套的torch.no_grad()上下文管理器。它的作用是关闭梯度计算推理时不仅省内存速度也会快不少。我见过不少初学者的推理代码没有这个结果一个batch就把显存吃满了。正确的推理范式就是上面这段eval()no_grad()。5. 深度踩坑nn.Module使用中的常见问题与排查5.1 参数不在parameters()里——老生常谈但每次都有人栽这个坑前面提过一次但我觉得值得单独拎出来再强调一遍。以下三种写法都可能导致模型参数无法被优化器扫描到# 错误1直接用普通list存子模块 self.layers [nn.Linear(64, 64) for _ in range(3)] # 错误2用普通Tensor代替nn.Parameter self.weight torch.randn(64, 64) # 错误3把子模块放进普通dict self.blocks {enc: nn.Linear(64, 64)}这三种情况的共同点是Module的子模块和参数虽然被赋值给了实例属性但没有走nn.Module的注册机制因此不在model.parameters()里面。排查方法很简单param_count sum(p.numel() for p in model.parameters()) print(param_count) # 如果你觉得模型应该有几十万参数打印出来只有几百那肯定有模块没被注册或者用model.state_dict()的键名来排查。正常注册的键名会像layers.0.weight、blocks.enc.weight而漏注册的模块的值根本不会出现在state_dict里。5.2inplace操作引发的梯度问题很多激活函数都有一个inplaceTrue的参数比如nn.ReLU(inplaceTrue)。它的作用是直接在输入Tensor上做修改省一点内存。但inplace操作有个隐藏风险如果输入Tensor是一个需要梯度的叶子变量并且它在后面的计算中还会被用到inplace修改可能破坏反向传播需要的计算图节点导致报错RuntimeError: a leaf Variable that requires grad is being used in an in-place operation.我建议在自定义Module的forward里尽量别用inplace特别是当输入x来自上层模块的输出而非模型入口时。虽然nn.ReLU(inplaceTrue)在ResNet这种深层模型里能省不少内存但为了稳妥你还不如用F.relu(x)让PyTorch自动管理内存。5.3 设备不匹配Expected all tensors to be on the same device这个问题几乎每个写PyTorch的人都遇到过。模型在GPU上输入是CPU Tensor一 forward 就报这个错。更隐蔽的情况是你自己写的Module里有一个常量Tensor是在CPU上创建的比如class MyModule(nn.Module): def __init__(self): super().__init__() self.scale torch.tensor(2.0) # 在CPU上且不是Parameter也不是buffer def forward(self, x): return x * self.scale当x在GPU上时x * self.scale就会爆设备不匹配错误。解决办法有两个一是把这个常量注册成buffer并且在初始化时不指定设备然后用model.to(device)统一移动二是在forward里动态获取x的设备def forward(self, x): return x * self.scale.to(x.device)第二种方法更省事但每次前向都调用一次.to()会有微小开销。更推荐的做法是把所有需要持久化的状态都用register_buffer注册这样model.to(device)会自动把它们一起搬走。5.4 钩子的使用当你想拿到中间层输出时nn.Module的钩子机制是个宝藏但很多人不知道。举个例子你想拿到模型中间层的输出来做可视化但又不想改原模型的forward代码这时候可以注册一个forward hookfeatures {} def hook_fn(module, input, output): features[mid] output model nn.Sequential( nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 10), ) handle model[1].register_forward_hook(hook_fn) x torch.randn(1, 784) model(x) print(features[mid].shape) # torch.Size([1, 256]) # 用完之后记得移除钩子 handle.remove()钩子函数接收三个参数模块本身、输入、输出。在hook里你可以对输出做任意操作比如打印shape、保存特征、甚至修改梯度。需要注意的是forward hook是在forward返回之后立刻被调用的所以拿到的是该层原始输出还没经过后面层。还有register_backward_hook可以拿到梯度但那个API在PyTorch版本间变更比较大不同版本签名不太一样用的时候最好查一下当前版本的文档。我在实际项目里用得更做的是自定义backward或者在loss层面做梯度修改backward hook用得相对少。5.5 打印模型结构的心得我调试模型的时候最常用的就是print(model)或者print(model.state_dict().keys())。前者能看到嵌套的模块结构后者能看到参数名。结合这两者基本上能做到看一眼就知道模型哪一层哪一维不对。另外如果是超深模型比如有几百层的Transformerprint(model)会把整个结构刷屏刷到看不见。这时候可以用model.named_modules()来做更精细的检查比如打印出所有Linear层的维度for name, module in model.named_modules(): if isinstance(module, nn.Linear): print(f{name}: in{module.in_features}, out{module.out_features})这个技巧在做模型结构对比或者解析第三方开源模型时特别实用。6. 进阶玩法利用Module信息做模型剪枝与稀疏化nn.Module的参数管理系统还有一个大用途模型剪枝。所谓剪枝就是去掉那些对最终输出贡献很小的参数让模型稀疏化从而压缩体积、加速推理。PyTorch官方在torch.nn.utils.prune里提供了相关工具但它也依赖nn.Module的参数注册机制。简单来说剪枝就是对某个nn.Module底下的weight做mask让部分参数变成0并重新注册新的参数。因为Module的参数管理是基于_parameters字典的剪枝模块可以直接替换掉已有的Parameter。import torch.nn.utils.prune as prune layer nn.Linear(64, 64) prune.random_unstructured(layer, nameweight, amount0.3) # 被剪掉的参数以mask形式存在 print(layer.weight_mask)这里的原理是prune模块会把原始的weight移除存到weight_orig里同时新增一个weight_mask前向传播时实际用的是weight_orig * weight_mask。这只能在nn.Module上实现换了其他自制的类很难做得这么干净。这种操作在实际部署场景中使用价值很大。比如你的模型有10M参数用30%稀疏化之后存储体积直接下降30%配合专门的稀疏推理库还能换来实打实的速度提升。当然剪枝本身是一门很深的学问我这里只是告诉你所有这一切的地基依然是nn.Module那套参数注册与重写的机制。7. 我的实践经验总结写了这么多年模型如果让我说nn.Module最值得记住的一句话那就是它是PyTorch所有网络行为的总调度中心。你写forward只是定义了数据怎么流而参数管理、模式切换、设备迁移、钩子机制、序列化保存这些全都是nn.Module在背后默默处理的。理解了这一点很多看似玄学的报错其实都能找到方向。给你几条我自己的实操建议算是我踩坑多年换来的经验第一所有子模块和参数一律通过nn.Parameter、register_buffer、nn.ModuleList、nn.ModuleDict来管理不要用Python原生容器。这能省掉你后面至少十个小时的debug时间。第二自定义Module时把权重初始化写成一个独立方法并在__init__里显式调用。虽然PyTorch默认初始化通常能跑但显式初始化让你的模型行为可预期复现实验的时候会省很多事。第三保存模型只存state_dict加载时善用strictFalse。迁移学习场景下这个习惯能让你少写很多模型结构对齐的脏代码。第四如果模型行为不符合预期先检查model.train()/model.eval()模式是否正确再检查参数是否在parameters()里面最后才去怀疑数据问题。这三个地方能覆盖90%的loss不下降推理结果不稳定问题。nn.Module的用法本身不难难的是建立起对整个参数管理体系的全局认知。当你真的把它吃透了后面再看nn.Transformer、nn.LSTM、各种自定义Layer都会觉得豁然开朗——因为所有这些东西本质上都是一个个注册了参数、实现了forward的Module在组合叠加。
