ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

从三行训练循环拆解PyTorch优化器参数更新原理与实战

2026/9/28 5:51:46 拓冰建站 浏览量
从三行训练循环拆解PyTorch优化器参数更新原理与实战 从 PyTorch 安装好、到第一次跑出 loss 下降很多人卡在的不是模型结构而是那条看起来只有三行的训练循环optimizer.zero_grad()、loss.backward()、optimizer.step()。看教程时大家都会说“这三行就是优化器在更新参数”可一旦自己调模型就会发现梯度爆炸、loss 不降、不同 batch 之间结果乱跳问题全都出在这三行背后的细节里。这篇就专门拆 PyTorch 优化器参数更新的完整步骤从底层 API 到手写 optimizer把“参数到底是怎么被改的”这件事讲透。适合正在学 PyTorch 的中级同学也适合跑过几个项目但一直没深究过 optimizer 实现的同学。如果你正处于“能跑通、但说不清为什么”的阶段这篇文章应该能帮你把最后那层窗户纸捅破。1. 优化器的本质训练循环里那个看不见的推手1.1 一套训练循环到底在跑什么先看一个 PyTorch 最典型的训练步optimizer.zero_grad() loss loss_fn(model(x), y) loss.backward() optimizer.step()很多人以为“训练”就是从 model 前向得到输出再算 loss但这只是前半场。真正让模型发生变化的是最后一行step()。前向传播只是算了一个“预测值”loss 只是给这个预测值打个分backward()把“打分结果”转化成每个参数的梯度而step()才拿着这些梯度去修正参数。把这件事拆开看完整链路其实是optimizer 在初始化时拿到model.parameters()的引用它知道该改哪些参数。每个参数 tensor 上都会有一个.grad属性backward()负责往这个属性里填梯度。step()遍历所有参数根据.grad和自身维护的状态比如动量计算一个更新量然后原地修改参数的数据。这里最关键的一点是update 是原地操作。PyTorch 的优化器不会重新创建参数而是用param.data.add_或者param.add_这种原地方式把新值写回原 tensor同时会记录param._version的增长用来配合反向传播的 graph 校验。这也是为什么你在训练时如果用param.data直接改参数会导致 autograd 的版本检测报错——它发现“这个张量在我不知情的情况下被改了”。1.2 为什么说 step() 之前先 zero_grad() 是必须的初学者最常见的困惑就是为什么每次都要 zero_grad梯度不是 backward 自动算出来的吗为什么不清零就会累积原因是tensor.grad是累加式的。PyTorch 的backward()默认把新梯度加到已有的.grad上而不是覆盖。这个设计本来是为了支持 gradient accumulation——先在多个小 batch 上分别 backward攒够一批梯度再 step。但如果你的训练循环里不调zero_grad()每个 step 的梯度就会不断叠加参数更新的方向直接被历史梯度带偏loss 曲线会变得非常诡异。我实际调模型时遇过一次很隐蔽的情况模型里有一个分支在 loss 为 0 时跳过 backward而我没在分支里单独处理 zero_grad结果那个分支的参数永远用的是上一次的梯度训练完全混乱。后来养成的习惯是zero_grad 的位置必须放在逐参数梯度的最早使用点之前也就是每个 step 的最开头而不是.step()之后。两者在大多数情况下都能跑但放到最开头才是语义最干净的写法。另外zero_grad 还有两个等价写法要了解。optimizer.zero_grad()内部其实调用了model.zero_grad()——把模型里的所有参数的.grad置 None而不是置 0。置 None 比置 0 好因为 PyTorch 底层在 backward 时如果发现 grad 是 None可以直接为它新建内存不用先清零再写一遍稍微省一点开销。你如果追求极致性能也可以用optimizer.zero_grad(set_to_noneTrue)这是 PyTorch 1.7 之后默认的行为。不同 PyTorch 版本之间的这个默认差异也常常影响老代码迁移后的行为值得注意。2. PyTorch 优化器 API 的底层逻辑2.1 先搞清楚 optimizer 初始化时你传入的是什么optimizer torch.optim.SGD(model.parameters(), lr0.01, momentum0.9)这个构造函数内部做了三件事把 params 解析成 param_groups、把超参存进 defaults、给每个参数分配专属状态容器 state。param_groups 是一个 list of dict每个 dict 里有params参数列表、lr、momentum等字段。你构造一个 optimizer其实是在维护一组“参数块和它们各自的超参映射”。默认情况下你传入的所有参数都在同一个组里。但如果你想要不同的层用不同学习率就可以手动构造多个组optimizer torch.optim.SGD([ {params: model.backbone.parameters(), lr: 1e-4}, {params: model.head.parameters(), lr: 1e-3} ], lr1e-2)注意这里有个坑你在 list 里给某个组指定了 lr它就会覆盖 defaults 里的全局 lr。defaults 里的值只是兜底。很多人写代码时以为“外层 lr1e-2 会统一覆盖”结果发现 backbone 一直按 1e-4 在训这就是对 param_groups 层级关系不够熟。理解 param_groups 对后面排查问题很重要。因为 optimizer 打印出来就是打印这些 group调参、保存 checkpoint、恢复训练时都要从这个结构里取 lr 和动量。2.2 step() 内部到底做了什么拿最简单的 SGD不带 momentum举例它 step 的核心逻辑只有两个层次外层遍历 param_groups内层遍历每个 group 里的参数。for group in self.param_groups: for p in group[params]: if p.grad is None: continue d_p p.grad.data p.data.add_(-group[lr], d_p)先忽略 weight decay 和 momentum这个循环的本质就是每个参数减掉 学习率 × 它的梯度。add_(-lr, d_p)做的是p p - lr * d_p。你可能会问为什么梯度不是加而是减因为损失函数求导得到的梯度方向是 loss 上升最快的方向我们要让 loss 下降就得往相反方向走所以是减。这个“减”字是所有优化器的原点。带 momentum 的 SGD 只是把 d_p 先替换成一个“历史梯度的指数移动平均”Adam 只是在这个基础上再做自适应缩放。后面第 3 节会具体展开。另外一个细节if p.grad is None: continue。这个判断非常重要。如果你的模型里有的参数不需要梯度比如某些冻结的层、或者 requires_gradFalse 的 embeddingbackward 不会给它生成.gradstep 必须跳过它否则直接访问p.grad.data会崩。这也是为什么很多教程里说“冻结层时 optimizer 可以不管”因为 step 自己会跳过这些参数。3. 手写一个 50 行的 optimizer彻底看透参数更新公式3.1 从最朴素的 SGD 开始理解一个东西的最好方式是自己实现一遍。我建议任何学过 PyTorch 的人都至少手写一次 optimizer不用写多完整能把核心更新逻辑跑通就足够了。下面这个类其实就是一个可用的 SGD不带动量import torch class BareSGD: def __init__(self, params, lr0.01): self.params list(params) self.lr lr def zero_grad(self): for p in self.params: if p.grad is not None: p.grad None def step(self): for p in self.params: if p.grad is None: continue p.data.add_(-self.lr, p.grad.data)用这个类你会发现训练效果和torch.optim.SGD(lr0.01)完全一样。为什么因为官方 SGD 在没有 weight decay、没有 momentum 的情况下的确就是这两行。从这里开始任何对“优化器更新参数”的疑问都能落到具体代码上而不是记忆概念。把p.data.add_(-self.lr, p.grad.data)展开其实等价于p.data p.data - self.lr * p.grad.data区别在于前者是 inplace 操作不会创建一个新的 tensor 覆盖param.data后者会先算出一个新 tensor再赋给p.data。注意不要直接写p p - self.lr * p.grad这会断开 p 和原参数的引用关系optimizer 后续再也找不到这个参数了。这是手写优化器最容易犯的错。3.2 加上动量从“只看当下”到“记住惯性”朴素 SGD 的问题是每次都只看当前梯度一旦梯度方向来回抖更新轨迹就会很震荡。动量momentum方案是给更新量加一个惯性把过去几步梯度的方向记下来别让每次更新都被最新一个梯度牵着鼻子走。class MomentumSGD: def __init__(self, params, lr0.01, momentum0.9): self.params list(params) self.lr lr self.momentum momentum self.buffer {p: torch.zeros_like(p.data) for p in self.params} def step(self): for p in self.params: if p.grad is None: continue buf self.buffer[p] buf.mul_(self.momentum).add_(p.grad.data) p.data.add_(-self.lr, buf)这里 buffer 就是官方 SGD 里的 momentum_buffer。buf momentum * buf grad是梯度的指数移动平均把历史的梯度方向记下来。如果过去几次梯度都指向同一个方向这个累计值会越来越大更新步幅会被放大如果梯度来回反向buffer 会被磨平步幅被压缩。这就是为什么动量 SGD 在深度学习里比朴素 SGD 稳得多。注意 PyTorch 官方的带动量 SGD 默认还带一个 nesterov 选项Nesterov 会让梯度在“当前速度指向的预计位置”再算一次实现上比普通的动量形式复杂一点。实际训练里我很少开 nesterov默认 momentum0.9 就够用了除非你在跑图像分类类任务想要多一点点收敛加速。3.3 Adam二阶矩自适应缩放Adam 比 SGDmomentum 多维护了一个“梯度平方的指数移动平均” v。用 v 去归一化学习率让更新量在每个参数维度上量级相近class BareAdam: def __init__(self, params, lr1e-3, betas(0.9, 0.999), eps1e-8): self.params list(params) self.lr lr self.betas betas self.eps eps self.m {p: torch.zeros_like(p.data) for p in self.params} self.v {p: torch.zeros_like(p.data) for p in self.params} self.t 0 def step(self): self.t 1 beta1, beta2 self.betas bias_corr1 1 - beta1 ** self.t bias_corr2 1 - beta2 ** self.t for p in self.params: if p.grad is None: continue grad p.grad.data self.m[p].mul_(beta1).add_(1 - beta1, grad) self.v[p].mul_(beta2).add_(1 - beta2, grad * grad) m_hat self.m[p] / bias_corr1 v_hat self.v[p] / bias_corr2 p.data.add_(-self.lr, m_hat / (v_hat.sqrt() self.eps))这里有个经常被忽略的点self.t是全局步数计数。Adam 的偏差修正bias correction依赖它。如果你从头训练时都不修正前几步的 m 和 v 都是从 0 开始累计会明显偏小导致前几步参数被修得过大加了修正后前几步的步长才是合理的。PyTorch 内部对 Adam 也是这么实现的。v.sqrt()在实现时用v_hat.sqrt()而不是torch.sqrt(v_hat)其实完全等价只是写法不同。你手写时注意在分母上要加 eps不要直接把 eps 加到 v 上否则 v 很小的维度会直接曝掉。加在分母里是为了数值稳定性。这一段看完你应该能回答一个问题了Adam 更新的本质是“每个参数维度的学习率都不相同”而 SGD 是整个模型共用同一个学习率。这也是为什么调 Adam 时全局 lr 通常可以在 1e-3 附近而 SGD 经常要从 1e-2 起调。不是玄学是数学结构决定的。4. 参数更新自己动手时的常见坑4.1 你的 loss 不降可能只是梯度方向错了step()本身是个很机械的操作它只会按照你给它的梯度去移动参数。如果 loss 不降问题大多出在“梯度算错了”或者“更新方向反了”。每次调试时我都提醒自己optimizer 不会替你做任何判断它的“职责”就是把.grad里的数值翻译成参数位移翻译系数是 lr。所以当模型不收敛时把怀疑对象锁定在梯度质量上往往比盯着模型结构效率高很多。一个非常常见的低级错误写自定义 loss 时把符号写反。比如你想最大化某个 score于是直接loss -score然后正常 minimization 训练这没问题。但如果你在代码里本意是loss score却在 backward 前又手动对 loss 取了反或者是自己实现了输出层的偏导数但符号弄错最后表现为loss 曲线一路向上、或者反复横跳。排查方法很土但有效在训练第一个 batch 前手算一个简单输入的输出和梯度对比一下你理解的传播过程。我通常会在模型刚初始化时选一个固定样本把model(x)打印出来手动心算一遍 forward 的第一层看输出是否对得上。如果对不上先别调 lr先把你的模型 forward 和 loss 计算逐层 debug 一遍。还有一类问题跟学习率相关。lr 太大loss 会直接发散到 inflr 太小loss 下降慢到看起来像没在训。最直接的诊断是第一步 step 以后 loss 有没有轻微下降。如果第一步反而暴增通常就是 lr 太大了。你可以从 3e-4 这类保守值开始逐步往上试或者用 learning rate finder。4.2 梯度累积一个必须理解的用法前面说 zero_grad 的“累加特性”可以用于梯度累积。场景是你想用 64 的 batch size但显存只放得下 16。你可以每 16 个样本算一次 loss 和 backward但不 step攒满 4 次后再 stepaccumulation_steps 4 for i, (inputs, targets) in enumerate(dataloader): outputs model(inputs) loss loss_fn(outputs, targets) loss loss / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这里有个细节要把 loss 除以 accumulation_steps 再做 backward。因为不除的话4 次累计梯度相当于 batch 256 的梯度量级而你 lr 还是按 batch 64 调的更新步幅就会被放大 4 倍。这是最容易被忽略的“参数更新量变化”来源。注意梯度累积时一定要把 loss 除以累计步数再做 backward否则参数更新量会偏离预期。另外注意optimizer.step()和zero_grad()在这里的顺序先 step 再 zero_grad。因为本轮参数已经用累计梯度更新了下一轮的梯度必须清零。如果你把 zero_grad 放到了 if 外面每步都清梯度累积就没意义了。4.3 weight decay 到底加在哪一步PyTorch 的 SGD 里的 weight_decay 默认实现是在梯度上直接加weight_decay * param也就是把 L2 正则化合并到梯度里再一起做动量/更新。但 Adam 不一样PyTorch 的 Adam 默认走的是 decoupled weight decay也就是在最终更新时直接按比例把参数拉向 0而不是把 weight decay 项塞进梯度里。这两者的差别普通训练里很难看出来但在预训练、大规模微调里影响明显。PyTorch 也单独提供了torch.optim.AdamW它就是做了 decoupled weight decay 的 Adam。你如果是从别的框架迁移过来的注意不要用 Adam 直接传weight_decay1e-4就当 AdamW 用了公式不是同一个。我常用的经验图像分类、微调场景SGDmomentumweight_decay5e-4 是很经典的组合。大规模 Transformer 类模型AdamWweight_decay0.01 或 0.1 更稳。用了 weight_decay 之后如果 loss 出现奇怪振荡先检查你模型里有没有用 BatchNorm 且没有正确分组——通常 BatchNorm 和 bias 这两类参数不应该加 weight decay因为它们本身数量少、却对整个训练稳定性影响很大。4.4 用 lr_scheduler 时最容易出错的顺序很多人的认知里优化器和学习率调度器是两个独立的东西但在 PyTorch 的 API 设计里scheduler 需要握有 optimizer 的引用才能修改 param_groups 里的 lr。所以它们的 step 顺序有讲究不是随便调都行。正确写法是for epoch in range(num_epochs): for batch in dataloader: optimizer.zero_grad() loss ... loss.backward() optimizer.step() scheduler.step()注意scheduler.step()在 epoch 结束调用不是每个 batch 都调除非你用 OneCycleLR 这类 per-step 的 scheduler。而且scheduler.step()必须在optimizer.step()之后因为 scheduler 要读 optimizer 当前的 lr 才能做衰减。另一个坑恢复 checkpoint 时scheduler 的 state 也需要保存和恢复。很多人在中断训练后直接加载模型权重和 optimizer state却忘了把 scheduler 的 last_epoch 恢复回去导致恢复后的学习率跳到了初始状态。进阶做法是把三者一起存进同一个 dictcheckpoint { model: model.state_dict(), optimizer: optimizer.state_dict(), scheduler: scheduler.state_dict(), epoch: epoch, } torch.save(checkpoint, ckpt.pth)这样恢复时一次性 load 回来不容易漏。5. 把步子迈开从 LSTM 训练到模型导出5.1 一个完整的 LSTM 序列训练循环光讲参数更新公式可能还是有点抽象我拿一个实际的 LSTM 序列预测任务为例把整个训练循环嵌进去看。假设要做的是时间序列预测输入是 (batch, seq_len, input_size) 的窗口数据import torch import torch.nn as nn class SimpleLSTM(nn.Module): def __init__(self, input_size, hidden_size, num_layers): super().__init__() self.lstm nn.LSTM(input_size, hidden_size, num_layers, batch_firstTrue) self.fc nn.Linear(hidden_size, 1) def forward(self, x): out, _ self.lstm(x) # 取最后一个时间步的输出 return self.fc(out[:, -1, :]) model SimpleLSTM(input_size8, hidden_size32, num_layers2) optimizer torch.optim.Adam(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, patience5) loss_fn nn.MSELoss() for epoch in range(50): epoch_loss 0.0 for X, y in train_loader: optimizer.zero_grad() pred model(X) loss loss_fn(pred, y) loss.backward() # 可选梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() epoch_loss loss.item() * X.size(0) epoch_loss / len(train_loader.dataset) scheduler.step(epoch_loss) print(fepoch {epoch} loss {epoch_loss:.4f})几个点可以展开说一下。为什么要用out[:, -1, :]而不是out全部送入全连接层因为序列预测的常见做法是拿最后一个时间步的隐状态做输出如果输入输出对齐模式不同取法也不同但这个“取最后一个时间步”的做法在很多任务里是默认选择。为什么 LSTM 训练时要加clip_grad_norm_因为 LSTM 存在长时序梯度爆炸的问题梯度范数很容易变得非常大。clip_grad_norm_做的事情是计算所有参数的梯度 L2 范数如果超过 max_norm就把它们等比缩小回 max_norm。这个操作放在 backward 之后、step 之前。它不影响梯度方向只影响更新幅度。在序列模型里我基本默认加上max_norm 取 1.0 到 5.0 都能接受。还有一点LSTM 是带内部循环的结构RNN 系的模型训练时对设备环境比较敏感。如果你在 Windows 的 CPU 环境、还是在带 GPU 的环境跑速度差异可能非常大。这也是为什么很多人一开始装 PyTorch 时会卡在环境上——没有对应 CUDA 的 PyTorchGPU 用不上LSTM 这种循环结构在 CPU 上跑真的会让人怀疑人生。装 PyTorch 时先确认 python 版本和 PyTorch 版本的对应关系再用torch.cuda.is_available()验证环境是省时间的第一步。5.2 优化器状态对模型导出的影响训练完之后你要把模型部署出去一般会选择转成 ONNX。一个很多人忽略的点是导出模型时导出的是 model.eval() 状态下的 forward 计算图和 optimizer 一点关系都没有。但为什么我在这里特地说因为在我实际导出的经验里最常见的坑恰恰出在“模型里挂着训练专属的 buffer/状态”。举个例子如果你用了 BatchNorm训练时它的 running_mean/running_var 会随 forward 一起被更新。转 ONNX 之前必须确认model.eval()否则 running_mean 可能还在变导出后的模型行为不稳定。RNN/LSTM 则要注意初始隐状态h0, c0的处理方式导出时通常要能够接受外部的 init_state或者在模型内部固定初始为零向量。再看 ONNX 导出本身。PyTorch 里最常用的写法是dummy_input torch.randn(1, seq_len, input_size) torch.onnx.export( model, dummy_input, lstm.onnx, opset_version11, input_names[input], output_names[output], dynamic_axes{input: {0: batch, 1: seq}, output: {0: batch}} )这里 dynamic_axes 声明哪些维度是动态的。如果不声明导出后模型只能接受固定形状这在真实部署里基本不可用。opset_version 我一般选 11 或更高老版本对 LSTM 这类算子的支持不够完整。导出后建议用 onnxruntime 验证一遍输出和 PyTorch 的 eval 输出对比。差异通常很小但如果发现差很多先回去看 eval 模式和模型里的归一化层。我遇到过最诡异的 case 是 dropout 忘了设 eval导出后输出比训练时多了随机性对比结果一直在抖动。6. 实操中的问题排查与最终心得6.1 参数更新环节的典型报错速查这里把我在社区里见过、自己也踩过的问题整理成一张速查表按频率排现象原因处理方式step() 报 Trying to backward through the graph a second time同一张计算图被 backward 两次且 retain_graph 未设每轮 forward 只 backward 一次若必须重复用backward 里加 retain_graphTrueloss 从初始值一路飙到 inflr 过大或 loss 符号反了先调小 lr 到 1e-5 试跑检查 loss 表达式加载 checkpoint 后 loss 不像恢复训练的曲线optimizer.state_dict 或 scheduler.state 没恢复三种 state 一起 load注意 optimizer 必须和模型用同样结构重建某些层参数一直不动这些层没有梯度requires_gradFalse 或前向没经过它们打印参数字典检查 requires_grad确认 forward 是否真的用到这些层参数更新后梯度历史被破坏报 one of the variables needed for gradient computation has been modified在 backward 前对用于计算的 inplace tensor 做了修改不要在 backward 前用 inplace 操作改模型输入、隐藏状态这张表看起来很简单但每个问题背后都对应一个具体的“参数更新前置/后置条件”。比如最后一条我见过有人为了省显存在loss.backward()之前就对输入做了x x.view(-1)这种 inplace 修改结果 autograd 找不到原来的计算图节点。结果就是整个训练循环白跑。遇到问题时我习惯先做一个“最小化复现”把模型缩到一层线性层固定数据跑一次 backward 和 step看能不能稳定复现报错。能复现问题就锁在代码逻辑不能复现多半是数据或模型结构里的偶发问题。这个习惯帮我省了大量排查时间。6.2 我对优化器参数更新的几条实操心得写到这里有很多初学者问过我“你到底是怎么把优化器这块搞明白的”其实方法很笨把官方 optimizer 的源码打开一行行读再自己复现一版。torch.optim.SGD的源码大概只有一百多行你完全可以读得懂。读懂之后再回去看那三行训练循环你会觉得 PyTorch 的设计其实非常直接——zero_grad、backward、step对应的就是“清空旧账、计算新账、按账改参数”三个动作。还有一条经验是不要迷信优化器的黑盒调参。很多人一上来就开一堆自动化调参工具但我个人建议先把 optimizer 的状态打印出来看几遍。在训练中插入一行print(optimizer.param_groups[0][lr])检查 learning rate 在什么时候跳变这比任何可视化工具都能更快地帮你发现问题。最后给一个小技巧。如果某个模型训练非常不稳定你可以先用最朴素的SGD(lr1e-3)跑 50 步观察 loss 是否下降。如果用 SGD 能训但 Adam 不能说明自适应学习率被极端梯度带偏了这时去查你的数据里有没有异常值、loss 函数有没有除以过大的量纲。用 SGD 做“基线诊断”是我排查训练问题最常用的第一步。我自己回头看真正搞懂优化器的那天就是把torch.optim.SGD的源码打开、一行行读完的那天。从那时起再遇到训练不收敛我都会先回到 optimizer 这一步想想参数到底是怎么被改的而不是急着换模型结构、换损失函数。希望你下次在处理自己的项目时也能在这三行训练循环里看到比我当初更多的门道。