ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch torch.nn 底层逻辑:Parameter、Module 与梯度管理

PyTorch torch.nn 底层逻辑:Parameter、Module 与梯度管理 第一次用 PyTorch 跑通神经网络的人大概率都有过这样的经历对着 nn.Sequential 抄了一段 LeNetmodel(x) 输出 10 维 logitsloss.backward() 一把梭准确率还不错。代码确实能跑但 torch.nn 究竟做了什么可能很少有人仔细想过——为什么必须继承 nn.Module为什么权重放在 Parameter 里为什么 model(x) 不能随意替换成 model.forward(x)为什么某个参数的 grad 莫名其妙是 None。这篇文章想把 torch.nn 的这些底层逻辑串成一条线讲清楚它的设计哲学和实际坑位。这里不纠结安装和 CUDA 配置默认你已经在正常环境里跑通过 PyTorch 2.x。文章适合两类人一类是刚把模型跑起来、想从“抄代码”过渡到“自己设计网络”的初学者另一类是写过一些自定义层、但遇到参数注册、梯度调试、模型导出问题时会卡住的实践者。后面所有代码我都基于 PyTorch 2.x 写版本太老的接口差异自己留意一下。1. torch.nn 的定位它到底帮你封装了什么1.1 从手写梯度到层化封装torch.nn 解决的问题先回到没有框架的时代。想训练一个两层神经网络你要自己定义 W1、W2 两个权重矩阵写前向计算再手动推 loss 对 W1、W2 的梯度。推完你还得自己写更新逻辑把梯度和学习率乘起来再减到权重上。整个过程麻烦不说还特别容易在链式法则的某一环上算错。Autograd 的出现解决了一个大问题只要你把前向计算用张量运算写出来PyTorch 会按照计算图自动帮你算反向传播不再需要手推梯度公式。但光有自动微分还不够。一个稍微像样的网络参数动辄几十万甚至上百万个如果这些权重都散落在普通字典或者全局变量里初始化、保存、优化器调度根本无从下手。torch.nn 给出的统一答案是把“一组参数 一个前向计算”打包成 Module再用容器把 Module 组装成更大的网络。神经网络本质上就是“带参数的函数”Module 是对这个概念的一级封装它让你不再关心参数放在哪个 list 里而是用一套标准接口去管理它们。这也是为什么你会发现PyTorch 里几乎一切网络相关的概念都长成 Module全连接层是 Module卷积层是 Module损失函数还是 Module甚至一个完整的 Transformer 也是由一堆子 Module 组合出来的大 Module。理解了这个“层层组合”的模型后面看任何开源代码都不会晕。1.2 Parameter、Tensor、buffer三种对象的分工Module 里能存的对象不少但真正决定模型行为的是三种nn.Parameter、普通 Tensor、buffer。很多人混用它们然后被各种诡异 bug 折磨。我把三者放在一张表里对照着看对象被 parameters() 收集被 state_dict 保存requires_grad典型用途nn.Parameter是是默认 True可训练权重比如 Linear 的 weight普通 Tensor 赋值为属性否否需手动设置临时常量但更多时候是写错了register_buffer 注册否是默认 FalseBatchNorm 的 running_mean 等模型状态nn.Parameter 是 Tensor 的子类它和普通 Tensor 的关键区别在于注册机制。当你把某个 nn.Parameter 赋值给 Module 的属性时它会被自动收集到 .parameters() 迭代器和 state_dict 里优化器默认只认识这些参数。普通 Tensor 不会被收集意味着 optimizer 根本不会更新它。buffer 是另一类容易被忽略的对象。有些东西不是“训练参数”但确实是模型状态BatchNorm 里的 running_mean、running_varTransformer 里的位置编码都需要在保存模型时一起存下来。这类对象用 register_buffer 来注册它会进 state_dict但不会被优化器更新。如果你把位置编码这种东西直接赋成普通 Tensor保存模型时它就会丢恢复模型时 key 对不上一顿报错。注意Module 里到底有多少参数、多少 buffer不要靠猜。交互环境里直接跑sum(p.numel() for p in model.parameters())和你手算的参数量对一下能避免大量“训练半天 loss 不降”的经典问题。1.3 用 20 行代码写一个自己的 Linear理论说多了容易飘直接上代码。下面的 MyLinear 只保留 nn.Linear 最核心的行为包含 weight、bias 两个可训练参数forward 里调用 F.linear 完成计算。import math import torch import torch.nn as nn import torch.nn.functional as F class MyLinear(nn.Module): def __init__(self, in_features, out_features): super().__init__() # PyTorch 的 Linear 权重形状是 [out_features, in_features] self.weight nn.Parameter(torch.empty(out_features, in_features)) self.bias nn.Parameter(torch.empty(out_features)) self.reset_parameters() def reset_parameters(self): # 复刻 nn.Linear 的默认初始化 nn.init.kaiming_uniform_(self.weight, amath.sqrt(5)) fan_in, _ nn.init._calculate_fan_in_and_fan_out(self.weight) bound 1 / math.sqrt(fan_in) if fan_in 0 else 0 nn.init.uniform_(self.bias, -bound, bound) def forward(self, x): return F.linear(x, self.weight, self.bias)你直接调 MyLinear(5, 3) 就能看到输出形状。这个例子的核心意义在于Module 的 forward 里只需要声明“输入怎么变成输出”反向传播由 autograd 根据计算图自动补齐。你不需要在这里写任何梯度公式这比手写网络时代不知道轻松了多少倍。2. forward 不是你想调就能调Module 的执行链路与钩子2.1 model(x) 与 model.forward(x) 之间的差距很多人写习惯了 model(x)但不知道为什么不能直接调 model.forward(x)。其实 nn.Module 重写了call在真正调用 forward 之前会走 _call_impl里面做的事情包括触发各种全局状态检查、调用 forward pre-hook、执行 forward、再调用 forward hook。如果你直接调 model.forward(x)这些 hook 全都不触发。所以遇到这种情况就要先怀疑自己注册了 hook 却不生效先去查是不是用了 model.forward(x)。另一个实际影响是在 TorchScript 追踪、ONNX 导出、以及某些编译优化路径里框架期望统一走 model(x) 这个入口。你绕过它等于绕过了一套标准的生命周期管理。从设计角度看Module 把自己做成可调用对象还有一层深意它让层和模型拥有统一的接口。你在 forward 里调用 self.encoder(x)这个 encoder 不管是单个 Linear 还是一个包含几十层的大模型调用语法完全一样。这种统一性对构建复杂网络非常重要你可以自由地组合、替换、嵌套模块而不需要关心内部结构。2.2 用 hook 偷看每一层的输出调试神经网络时最常见的需求就是确认每一层输出的 shape 对不对。你当然可以在每个 forward 里手写 print但那样要改模型代码调完还要删。更好的方式是注册 forward hook不改模型就能打印输入输出形状def shape_hook(module, inputs, outputs): print(f{module.__class__.__name__}: in{tuple(inputs[0].shape)}, out{tuple(outputs.shape)}) model nn.Sequential( nn.Linear(28 * 28, 128), nn.ReLU(), nn.Linear(128, 10), ) for layer in model: layer.register_forward_hook(shape_hook) x torch.randn(4, 28 * 28) model(x)输出会一行一行地显示 shape 经过每层的变化。如果你想看反向传播时某一层收到的梯度可以用 register_full_backward_hook它能把梯度从后往前逐层打印出来在排查梯度消失问题时比打印参数梯度更直接。我用 hook 的经验是它能让你把注意力集中在“网络本身”而不是“为了调试而写的临时代码”。调试完如果不想保留直接 remove 掉模型代码干干净净。2.3 nn.ReLU 还是 F.relu有状态与无状态的分界这是每个 PyTorch 用户都会纠结的问题。torch.nn.ReLU 是一个 Module 类torch.nn.functional.relu 是一个函数。对于 ReLU 本身它没有可训练参数也没有需要保存的运行状态所以两者在数值上几乎可以互相替代。但官方示例里为什么常用 nn.ReLU()因为有状态模块必须用类的形式存在。最典型的是 Dropout训练时要随机丢弃推理时不能丢。nn.Dropout 会自己在 train 和 eval 模式之间切换行为。如果你用的是 F.dropout(x, p0.5)就得手动传一个 training 标志一旦漏传推理时的行为就是错的。BatchNorm 更明显它要维护 running_mean 和 running_var这些状态必须挂在 Module 上才能被 state_dict 保存。你没法用函数调用来保存状态。我的习惯是结构性层一律用 nn.比如 nn.Linear、nn.Conv2d、nn.BatchNorm1d、nn.Dropout纯粹为了省一次类实例化的小激活函数用 F.也可以。但团队项目里最好统一成类因为代码可读性和状态管理都会更好。这里顺便提一下参数共享。同一个 self.fc nn.Linear(...) 如果在 forward 里被调用两次这两次计算共享同一份权重反向时梯度会累加。这是 Module 灵活性的直接体现也是很多 RNN 和孪生网络的基础。理解这一点你就知道为什么 Module 不是一个简单的函数而是一个可复用的组件实例。3. 训练前的三件大事参数注册、初始化和优化器绑定3.1 参数没注册会怎样一个常见的静默错误有一类 bug 最让人头疼不报错但结果全错。参数没注册就是典型的静默错误。假设你写了一个线性模型图省事把权重写成普通 Tensorclass BadLinear(nn.Module): def __init__(self): super().__init__() self.w torch.randn(10, 5) # 错了这不是 Parameter self.b torch.zeros(5) def forward(self, x): return x self.w self.bforward 可以正常跑但你 print(list(model.parameters())) 会发现是空列表。于是 optimizer torch.optim.SGD(model.parameters(), lr0.01) 等于什么都没优化。这里的梯度和初始化都正常可更新根本不发生loss 自然不降。你查网络结构、查数据、查学习率都找不到问题最后才发现参数压根没进优化器。正确写法是 self.w nn.Parameter(torch.randn(10, 5))或者调用 self.register_parameter(w, nn.Parameter(...))。如果你有一批不希望被优化器更新的中间结果就注册成 buffer或者在 forward 里临时创建普通 Tensor不要塞进 parameters。这个点对刚接触自定义层的朋友尤其重要因为内置层都替你注册好了自己写层时才暴露出来。3.2 初始化不是随便 fillXavier 与 Kaiming 的差别权重初始化听起来是个小细节但影响极大。想象一个有 100 层全连接的网络如果每层权重初始化为方差较大的随机值前向信号的方差会逐层膨胀几十层后输出直接爆掉如果权重都特别小信号又会逐层缩到接近零。深度网络训练不起来初始化常常是第一道关。torch.nn.init 里的初始化方法可以粗略分成两类。一类是 Xavier/Glorot适合 Sigmoid、Tanh 这类关于零点对称的激活函数另一类是 Kaiming/He适合 ReLU 及其变体。区分标准很简单Xavier 假设激活值围绕零对称分布Kaiming 特意考虑了 ReLU 会把一半神经元置零的事实。你用 ReLU 却配 Xavier深层信号会被系统性衰减用 Kaiming 配 Tanh又容易出现震荡。激活函数推荐初始化备注Sigmoid / TanhXavierGlorot均匀或正态保持信号在对称区间内稳定ReLU / LeakyReLUKaimingHe均匀或正态考虑 ReLU 置零导致的方差变化无激活 / 线性输出方差适中的均匀分布避免信号指数级放大或消失从源码看nn.Linear 的默认初始化是 kaiming_uniform_bias 则会根据 fan_in 生成一个均匀分布的范围并不是固定填 0。自定义层如果在 reset_parameters 里全填零或乱填 uniform会出现深层网络训练不稳定的问题。我的建议是写自定义层时先把 reset_parameters 按“复刻内置层默认行为”的标准实现再根据实验慢慢调不要一开始就放飞。3.3 optimizer 如何与模型参数建立联系optimizer torch.optim.Adam(model.parameters(), lr0.001) 这行代码所有人都写过但背后有一个隐藏的坑model.parameters() 返回的是生成器optimizer 在构造时会把它消费成一个参数组列表。如果你在模型还没完全创建好、参数还没注册时就构建 optimizer这个优化器拿到的可能只是空集合。所以在写训练循环时顺序永远是“先搭完模型、再建优化器、再开始训练”。如果中途给模型添加了新参数旧优化器不会自动认识它们需要用 add_param_group 手动新增。这也解释了为什么不要在图省事时把 optimizer 写进模型类的init里。参数集合是动态的把它和时间线绑定在一起是最稳的。如果你怀疑优化器没拿到参数一行命令就能验证print(sum(p.numel() for p in model.parameters()), len(list(model.parameters())))再和预期参数量对比一下。我几乎每次调新模型都会跑一遍这一步成本极低收益极高能省掉大量“loss 不降”的排查时间。4. 组合网络的工程实践容器、常用模块与损失函数4.1 Sequential 够用但 ModuleList 更灵活刚写网络时大家习惯把 self.fc1、self.fc2、self.fc3 一个一个列出来然后手写 forward 里的调用链。层数少还好层数一多就是灾难。于是大家开始用 nn.Sequential它把“严格按顺序执行”这个语义直接封装成 forward。好处是代码短、适合直线结构坏处是你没法在中间自由插入分支或条件。nn.ModuleList 是另一个更灵活的容器。它只负责“收集一堆子模块”不定义执行顺序。执行顺序由你手写的 forward 决定。自由度更高代价是多写几行代码。我判断的标准很简单网络结构是纯直线用 Sequential网络里有分支、跳连、动态层数用 ModuleList 加 for 循环。这里有个非常容易踩的坑如果你用普通 Python list 存层比如 self.layers [nn.Linear(...) for ...]这些层的参数不会被 model.parameters() 收集。因为 Module 只认 ModuleList、ModuleDict 这类容器普通 list 里装着的子模块不会自动注册。很多人把层存在 Python list 里结果 optimizer 看不到参数训练完全不动查半天才发现是容器类型用错了。4.2 一个标准 CNN 和可变深度 MLP 的搭建拿 MNIST 手写数字识别当例子一个极简 CNN 用 Sequential 写出来长这样cnn nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(64 * 7 * 7, 10), )输入是 [batch, 1, 28, 28]经过两次 MaxPool2d(2) 后空间尺寸从 28 变 14、再变 7所以全连接层的 in_features 是 64 × 7 × 7。第一次写 CNN 的朋友经常把这里算错导致 Flatten 之后维数不匹配。建议先手动把特征图尺寸推一遍再写 Linear 的输入维度不要指望报错信息帮你定位它只会告诉你一个很长的矩阵乘维度错误。再看一个用 ModuleList 实现可变深度 MLP 的例子class FlexibleMLP(nn.Module): def __init__(self, sizes): super().__init__() self.layers nn.ModuleList() for i in range(len(sizes) - 1): self.layers.append(nn.Linear(sizes[i], sizes[i 1])) def forward(self, x): for layer in self.layers[:-1]: x torch.relu(layer(x)) return self.layers[-1](x)这个类的好处是 sizes 传 [784, 256, 128, 10] 就能生成三层网络想改成 5 层不用改类定义只改 sizes 就行。在超参数搜索里特别实用。注意最后那个 Linear 没有接激活函数这是分类网络输出的常见约定——原样输出 logits后面交给损失函数处理。搭好网络之后最小训练骨架是这样的model FlexibleMLP([784, 256, 128, 10]) opt torch.optim.Adam(model.parameters(), lr1e-3) loss_fn nn.CrossEntropyLoss() for epoch in range(3): for x, y in loader: out model(x) loss loss_fn(out, y) opt.zero_grad() loss.backward() opt.step()这个循环虽然简单但是所有 torch.nn 训练流程的最小公因数。后面的梯度调试、分布式训练、混合精度都不会脱离这个骨架太远。4.3 损失函数怎么选不要重复 Softmaxtorch.nn 里的损失函数也是 Module这一点常被忽略。最常用的三个我来列一下损失适用场景注意事项nn.MSELoss回归任务输出层不要加 Softmax对异常值敏感nn.CrossEntropyLoss多分类任务内部已包含 LogSoftmax NLLLoss模型直接输出 logitsnn.BCEWithLogitsLoss二分类 / 多标签分类内部已包含 Sigmoid不要手动再套一层我见过最多的问题就是把 nn.CrossEntropyLoss 和模型最后一层的 nn.Softmax 一起用相当于做了两次 softmax。前向计算看起来没啥问题但反向传播的梯度会被二次 softmax 扭曲训练要么极慢要么直接发散。记住CrossEntropyLoss 的输入是 logits不是概率。如果你确实需要概率做可视化在训练之外另算 torch.softmax(model(x), dim-1)不要在模型 forward 里加 Softmax 层。5. 梯度从哪里来反向传播的实际观测与自定义算子5.1 在训练循环里确认梯度真的还在训练模型的本质是让梯度驱动参数更新。如果一个参数根本收不到梯度那你只能瞎猜哪里出了问题。最简单的办法是在训练循环里每隔几步打印参数的梯度范数for name, param in model.named_parameters(): if param.grad is not None: print(name, param.grad.norm().item()) else: print(name, grad is None)在本地调试时这个输出能让你快速定位“哪一层梯度消失了”“哪一层梯度爆炸了”。如果发现某一层的梯度范数比相邻层小好几个数量级可以考虑初始化是否合适、激活是否饱和或者这一层根本没参与最终 loss 的计算。很多“loss 降不下去”的案例最终都是在这个输出里发现某一层永远是 None。5.2 参数 grad 为什么是 None这里集中说一下我见过最多的几类原因。第一requires_grad 为 False。你手动设置过 param.requires_grad_(False)或者加载预训练权重后冻结了某些层那么该参数不会保留梯度这是符合预期的。第二参数没有参与 loss 的计算图。网络里常有多个分支比如一个有辅助分类头的模型辅助分支的某个参数如果没被 loss 引用那它的梯度就是 None。你要检查 forward 里这个参数是不是真的喂进了主干路径。第三在 torch.no_grad() 上下文里做了计算。验证集评估时人们习惯包一层 no_grad但如果缩进没控制好后面的训练代码也落进了 no_grad 块backward 就会报“does not require grad”。这种问题靠肉眼查缩进就能解决但确实经常发生。第四调用了 detach()。detach 会隔断计算图梯度流到那里就断了。如果你在 forward 里为了某个目的对中间张量做了 detach后续所有依赖它的参数都收不到梯度。这是自定义网络里最隐蔽的坑之一。排查这类问题最快的路径是逐层收缩在可疑 Module 上注册 full backward hook看看它实际收到了什么再往前检查 forward 路径上有没有 detach 或 no_grad。不要一上来就怀疑优化器绝大多数情况是接到上一个环节就断了。5.3 自定义 autograd.Function什么时候需要它普通自定义层继承 nn.Module 就够因为 forward 里用的是 PyTorch 自带算子反向由 autograd 自动处理。但有些场景必须定义自己的反向传播你要实现一个融合算子节省内存某个数学操作没有现成 API或者你想写一个不可导函数的自定义梯度比如 straight-through estimator。这时候用 torch.autograd.Function。举一个最简单例子自定义 f(x)x^2 的正向和反向class MySquare(torch.autograd.Function): staticmethod def forward(ctx, x): ctx.save_for_backward(x) return x * x staticmethod def backward(ctx, grad_output): x, ctx.saved_tensors return grad_output * 2 * x在 Module 里调用它class MyLayer(nn.Module): def forward(self, x): return MySquare.apply(x)原理上ctx 负责在 forward 和 backward 之间传递信息save_for_backward 保存前向计算需要的中间量backward 接收上游传来的梯度 grad_output再乘上局部梯度完成链式法则的一环。需要注意的是backward 返回的 Tensor 数量和 forward 的输入数量一致没有梯度的输入返回 None。如果你返回多了一个或少了一个PyTorch 会直接报错。后面接触 LSTM、Transformer 里的 attention 算子或者图神经网络里的消息传递很多高性能实现都是“custom Function Module”的组合。这一块搞明白再看开源模型里的 fused kernel心里会有底很多。6. 从训练到落地的几个常见坑保存、设备、导出6.1 state_dict 的保存与恢复训练到后面总要把模型保存下来。最常规的只存参数方式torch.save(model.state_dict(), model.pt) # 恢复时 model MyModel() model.load_state_dict(torch.load(model.pt, map_locationcpu))这里有两个经典坑。第一torch.load 要加 map_location 参数把状态字典里的 tensor 设备统一到 CPU否则在 GPU 上保存、CPU 上加载时会报设备不匹配。第二strict 模式下 key 必须完全一致。如果你之前用 nn.DataParallel 包过模型state_dict 的每个 key 会多出 module. 前缀。恢复时要么用 model.module 取原始模型要么手动 strip 前缀否则 load_state_dict 会同时报 missing key 和 unexpected key。如果只是保存参数还不足以完整恢复训练现场。断点续训时通常还要保存 optimizer.state_dict()、epoch、scheduler 等信息。这里不展开但你要有这个意识state_dict 是整个模型状态的最小可迁移单位理解它对排错很有帮助。6.2 train/eval 切换与设备一致性模型实例创建后默认是 training 模式。推理前一定要调 model.eval()否则 Dropout 还在随机丢弃神经元BatchNorm 也还在用当前 batch 的统计量而不是 running stats结果会忽高忽低。大量“训练时准确率挺高推理时结果不对”的求助帖很多就是忘了这行代码。设备一致性是我见过的报错频率最高的问题之一。常见错误是模型和输入不在同一设备执行了 model model.to(cuda)但数据循环里没有 x x.to(device)前向计算就会报 “Expected all tensors to be on the same device”。更隐蔽的是 label 在 CPU 上、logits 在 GPU 上计算 loss 时同样会炸。我的习惯是训练循环开头定义 device并把每一个进模型的 tensor 显式 .to(device)。不要相信“我确认搬过去了”的记忆有时候就是因为某一次数据处理分支漏了。6.3 转 ONNX 时的 trace 陷阱模型上线部署常用的路径是把 PyTorch 模型导出成 ONNX再交给推理引擎。一个基本导出示例model.eval() dummy_input torch.randn(1, 1, 28, 28) torch.onnx.export( model, dummy_input, model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version17, )第一条导出前先 model.eval()避免 Dropout 被固化成一个随机节点或者 BN 使用 batch 统计量导致推理语义错误。第二条ONNX Export 走的是 trace 机制它会根据一次虚拟输入把模型执行过程“录”下来。如果你的 forward 里有数据相关的 Python 控制流比如if x.shape[0] 10这类判断trace 只保留第一次输入走的那条路径。新版 PyTorch 可以用 torch.onnx.export(..., dynamoTrue) 走符号追踪或者干脆把动态控制流改成张量运算让网络结构完全由输入 shape 驱动。第三条不要在 forward 里 .item() 或调用 Python 的 float() 来影响后续逻辑。这些值会被 trace 固化成常量输入换一个 batch 大小就失效。这也是 PyTorch 模型和部署模型差异最常见来源之一。第四条导出后用 onnx.checker 或者 onnxruntime 跑一遍验证输出。特别是动态 batch记得把 batch 维标出来否则 ONNX 默认按固定 batch 导出服务端只能一次推理一个样本。我在实际项目里还踩过一个专门的坑自定义 autograd.Function 通常没法直接 trace 导出除非你为它写 symbolic 规则把自定义算子映射到 ONNX 已有的算子。具体写法是给 Function 加一个 symbolic 静态方法。这一块属于进阶内容但如果你计划部署带自定义层的模型迟早会碰到。最后说一点个人体会torch.nn 给你的不是一个死板的“层库”而是一整套管理参数、定义前向计算、挂载调试钩子的框架。把 Parameter、Module、容器、autograd 这条链路理解透后面不管写 CNN、LSTM 还是图神经网络都是在同一个心智模型上叠加新概念。遇到问题的时候先查参数集合、再看计算图、再验梯度这个排查顺序能帮你省下大量时间。
返回列表