PyTorch参数更新机制与优化器实战:从梯度裁剪到混合精度训练
发布时间:2026/9/11 14:12:08 作者:尧图编辑部 阅读量:1,286

刚接触PyTorch的时候我一直在用一套“标准三连”跑训练loss.backward()、optimizer.step()、optimizer.zero_grad()。当时我以为参数更新就这么简单直到后来调模型时频繁遇到loss震荡、不收敛、显存爆炸才回去把torch.optim的源码和PyTorch的自动求导机制翻了个底朝天。这篇文章我就把PyTorch参数更新的完整链路、主流优化器的内部差异、以及实际训练里常用的梯度裁剪、梯度累积、冻结层更新、混合精度训练这些进阶玩法一次性讲透。内容面向正在用PyTorch做深度学习、但还没有系统梳理过参数更新机制的读者也适合日常调参卡壳、想真正搞懂每行代码背后发生了什么的朋友。1. 梯度从loss传到paramoptim.step()之前那几步的真正含义1.1 计算图不是模型是“求导路线图”很多初学者会把model本身和PyTorch构建的计算图混在一起这是第一个需要纠正的概念。nn.Module只是把网络层以Python对象的形式组织起来真正用来求梯度的是前向传播过程中自动创建的计算图。当你调用loss criterion(model(x), y)时PyTorch会在背后记录所有张量运算的调用关系输入x怎么经过卷积、激活、全连接最后算出loss。这个运算链路就是一张有向无环图每个节点是一个Tensor每条边表示“这个Tensor是由另一个Tensor经过什么运算得到的”。Python的Tensor对象里有一个grad_fn属性它指向的是创建当前张量的运算节点。比如z x y那么z.grad_fn就指向AddBackward这个运算节点节点内部保存了参与运算的x和y的引用。这张图只有在对requires_gradTrue的Tensor做运算时才会自动构建。模型里的参数在初始化时默认requires_gradTrue所以前向传播一定会生成计算图。但如果你推理时用了torch.no_grad()计算图就不会被保存这既省内存也省时间。1.2 backward()累加梯度与zero_grad的先后约定调用loss.backward()时PyTorch从loss这个Tensor出发沿着计算图反向遍历对每个叶子节点通常是模型参数计算梯度并把结果累加到对应参数Tensor的.grad属性上。这里的关键词是“累加”。每执行一次backward().grad不是被覆盖而是在原有值上做加法。为什么这么设计因为PyTorch要支持梯度累积的用法显存不够时可以把一个大batch拆成几个小batch每个小batch分别算loss和backward梯度会累加在一起最后统一走一次参数更新效果等同于用大batch更新。正因为梯度是累加的所以在一次完整的参数更新完成后你必须手动把.grad清零否则下一轮backward时梯度会和旧梯度叠加更新方向就乱了。标准写法的顺序是optimizer.zero_grad() loss.backward() optimizer.step()zero_grad()必须在backward()之前。有人会写成loss.backward()之后、optimizer.step()之前调用zero_grad()这会导致当前这轮的梯度先被清零参数根本没吃到梯度就更新了算是新手比较容易踩的坑。1.3 set_to_noneTrue比“全填0”快不少optimizer.zero_grad()内部有两种清空策略。默认情况下是把.grad里的元素全部填成0但如果设置zero_grad(set_to_noneTrue)PyTorch会直接把.grad置为None。两者效果不一样吗绝大多数场景下是一样的因为backward()执行时只要requires_grad为TruePyTorch会为参数重新分配一个空的grad Tensor然后往里累加梯度。set_to_noneTrue的性能优势在于把grad置为None不需要遍历Tensor逐元素写0少一次内存写入操作训练速度会有微小但可感知的提升。我自己在训练ResNet这类大模型时把这一处改掉后每个iteration能省下几毫秒到十几毫秒积少成多一天训练下来省的时间相当可观。不过要注意置为None后如果你在step()前后自己写了读取param.grad的逻辑比如想打印梯度范数就得先判空否则会报AttributeError。实际开发里为了稳定我通常还是用默认的false只有在大规模训练中对性能极敏感时才开set_to_none。2. 主流优化器内部差异SGD、Adam、AdamW为什么要设不同更新参数2.1 动量的物理图像一个纯SGD更新长这样param - lr * grad问题是如果loss landscape在某些方向上很陡、在另一些方向上很平缓纯SGD会来回震荡收敛很慢。动量Momentum就是给更新引入“速度”的概念速度向量的更新公式是v momentum * v grad然后再用速度去更新参数。物理上理解就是一个小球在loss曲面上滚动累积的速度让它在平坦区域也不会停下来还能冲过局部极小值附近的沟壑。PyTorch里用torch.optim.SGD(model.parameters(), lr0.01, momentum0.9)开启动量默认momentum0。动量值一般设在0.8到0.99之间0.9是使用最多的默认值。调大动量会让收敛更平滑但也更容易冲过头训练后期可能需要配合学习率衰减来控制。2.2 从Adagrad到RMSProp再到Adam动量为每个参数使用统一的更新强度但真实网络的梯度尺度差异很大某些层的梯度数量级很小某些层梯度很大。针对每个参数自适应调整学习率的想法由此而来。Adagrad的思路是累积所有历史梯度的平方和作为分母对每个参数单独缩放学习率。问题是分母不断增大学习率会被压到趋近于0训练基本走不动。RMSProp的改进是用指数移动平均来替代历史平方和只关心最近的梯度尺度分母不会无限膨胀。RMSProp的更新公式里有一个eps防止除零PyTorch默认是1e-8。Adam相当于“带动量的RMSProp”再加上一阶矩估计。它维护两个状态一阶矩m梯度的指数移动平均和二阶矩v梯度平方的指数移动平均然后做偏差校正。日常用torch.optim.Adam(lr1e-3)作为默认起点这个学习率在绝大多数CV和NLP任务上都表现得很稳。三个优化器的更新差异用一句话概括SGD只认当前梯度RMSProp只看梯度尺度的近期波动Adam既看方向又看尺度所以它对初始学习率的敏感性更低。优化器维护状态主要特点典型学习率SGD动量速度收敛稳定泛化性好需要精细调参0.01~0.1RMSProp梯度平方的移动平均自适应学习率适合非平稳目标0.001~0.01Adam一阶矩二阶矩适应性强开箱即用0.0001~0.0012.3 AdamW的权重衰减与L2正则化不是一回事在PyTorch中Adam优化器里的weight_decay参数默认实现的是L2正则化在loss上加上所有参数的平方和乘以系数这样梯度里会附带一项weight_decay * param。但Adam维护的二阶矩会把正则项和真实梯度的量级混在一起做归一化结果导致权重衰减的实际强度被削弱了而且强弱还随梯度的变化而波动。AdamW把weight decay从梯度计算中拿出来改为在参数更新时直接对参数本身做缩放param - lr * (weight_decay * param)。这个过程发生在Adam更新之前不影响Adam内部的一二阶矩估计每个batch对参数的“拉向零”力度是固定的所以叫“解耦权重衰减”。因此两种优化器即使设置同样的weight_decay实际效果差异很大。我在训练Transformer类模型时用AdamW(lr1e-4, weight_decay0.01)比Adam(lr1e-4, weight_decay0.01)的收敛速度和最终效果都要好不少。如果你从别人代码里看到AdamW注意别把它当成加了正则的普通Adam来理解这两个在源码层面计算路径就不同。3. 训练不稳、显存不够时的更新策略梯度裁剪与梯度累积的正确姿势3.1 梯度裁剪救日志里的nan和loss spike训练过程中常常会遇到一种情况loss在某个iteration突然飙到好几千然后后面全是nan。常见原因之一是梯度爆炸尤其是深层的RNN、Transformer或加了较大的初始学习率时。梯度裁剪的思路很简单如果梯度的范数超过某个阈值就把所有参数的梯度按比例缩小让梯度总范数不超过阈值。PyTorch提供了两种方式# 按参数梯度的总范数裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 按每个参数的梯度绝对最大值裁剪 torch.nn.utils.clip_grad_value_(model.parameters(), clip_value0.5)日常用得最多的是clip_grad_norm_。它的计算过程把每个参数的梯度张量展平算所有梯度的全局二范数如果超过max_norm就用max_norm / global_norm这个系数依次缩放每个参数的梯度。注意一点clip_grad_norm_必须放在backward()之后、step()之前。很多人把顺序搞反在backward之前就裁剪这时梯度的值还是None或旧值裁剪等于白做。max_norm怎么选我通常在1.0附近起步如果是GAN或RNN这类容易不稳定的模型可以试0.1到0.5。太小会严重抑制真实的学习信号太大又起不到保护作用。从log里看如果裁剪前后梯度范数差异非常大说明梯度一直在撞击上限应该考虑降低学习率而不是无限调大max_norm。3.2 梯度累积把batch size撑大N倍的更新方法显存有限时想用更大的batch size训练模型一个通用做法是梯度累积。思路是不每个mini-batch都更新参数而是连续跑N个小batch累加梯度每N次backward后才执行一次step。实现方式有两种。第一种是手动控制accumulation_steps 4 for i, (inputs, labels) in enumerate(dataloader): outputs model(inputs) loss criterion(outputs, labels) loss loss / accumulation_steps # 把loss均分 loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这里一定要做loss / accumulation_steps否则累加N次后的梯度过相当于用了N倍学习率参数更新步长被放大训练很容易震荡。还有一种是利用PyTorch梯度累加机制配合with torch.no_grad()和detach去模拟但手动控制法已经够用代码也直观。要注意的坑是如果你的模型里有BatchNorm层梯度累积下只有最后一次backward之后会触发step但BatchNorm的统计量在每个mini-batch的前向里都会更新这与真正使用大batch时的统计更接近是有利的。可是如果累积步数太大BN的统计量会因为每个mini-batch的数据分布差异大而抖动这时需要同步BN统计量torch.nn.SyncBatchNorm.convert_sync_batchnorm。3.3 梯度累积和梯度裁剪连用的执行顺序梯度累积场景下用裁剪最稳妥的顺序是在每个step前对累积后的梯度做clip_grad_norm_。也就是说裁剪只需要在optimizer.step()之前执行一次而不是每个mini-batch都裁剪。如果每个小batch都裁剪会把还没累积完的梯度强行限制范围最终累积结果就不是真实梯度的平均更新方向和用大batch训练时差异会很大。正确的完整流程是for i, (inputs, labels) in enumerate(dataloader): outputs model(inputs) loss criterion(outputs, labels) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() optimizer.zero_grad()这条链路我自己在显存只有12G的卡上训练一批语义分割模型时验证过效果很稳。还有一个容易被忽略的问题如果数据集长度不能被accumulation_steps整除最后一个epoch可能有残留梯度没有step就进入了下一轮迭代需要自定义一下len(dataloader)或提前丢弃最后不足一个累积块的数据。4. 迁移学习和模型微调中的参数更新冻结层、分组学习率与BN的坑4.1 requires_gradFalse只是“不更新”不是“不计算”迁移学习很常见的操作是把预训练模型前面的层冻住只微调后面的分类头。参数冻结的写法很简单for param in model.parameters(): param.requires_grad False但我提醒一句requires_gradFalse只意味着这个参数在反向传播时不会被赋予梯度也不参与参数更新。但它的值在前向传播中仍然是参与计算的。而且由于PyTorch自动求导是按需创建计算图如果某个Tensor的requires_grad为False依赖它计算出来的中间结果若非叶子节点则整链路上的梯度都不会保留。所以冻结层并不会大幅减少前向和反向的计算量只是省掉了该层参数梯度计算和更新的开销。真正省显存和计算的方法是冻结后对该层输入做一次前向并detach()后续不再经过该层。比如把backbone的中间特征提前提取出来存到内存或磁盘再单独训练后面的分类器。这时才算彻底“不计算”。4.2 param_groups分组学习率微调时一个通用经验是预训练层用较小的学习率新加的随机初始化层用较大的学习率。PyTorch优化器里可以通过param_groups实现。optimizer torch.optim.Adam([ {params: model.backbone.parameters(), lr: 1e-5}, {params: model.classifier.parameters(), lr: 1e-3}, ], weight_decay1e-4)每个param_groups是一个字典可以独立设置lr、weight_decay、betas等超参数。optimizer.param_groups是一个列表索引对应你传入的每个分组修改分组学习率也能按组生效比如写optimizer.param_groups[0][lr] 5e-6。有个细节整个优化器创建时传入的默认学习率lr会被同步写进每个分组如果分组的字典里没显式设置lr就用全局的lr。写分组时容易漏掉分组里的lr字段导致两个分组的lr完全相同分组设置等于无效。4.3 冻结层时常被忽略的BN和Dropout状态微调时最隐蔽的坑在于BatchNorm层。即使你设置了requires_gradFalse冻结了base网络如果模型整体处于model.train()模式BatchNorm层仍会不断更新自身的running_mean和running_var同时计算当前batch的均值方差做归一化。这会导致冻结的训练阶段BN统计量被当前任务的图像分布不断覆盖原始预训练模型学到的数据分布信息逐渐丢失训练完成后推理效果反而变差。熟悉PyTorch的人知道model.eval()会把BN切到推理模式使用固定的running统计量同时Dropout层会失效。因此微调冻结backbone时更安全的是冻结整个backbone的BN为eval模式只把新加的分类头置于train模式实现上比较麻烦因为model.train()是全局生效。常见做法是把backbone里的每个子模块手动切换到eval这个可以通过自定义递归函数做到def set_bn_eval(module): if isinstance(module, torch.nn.modules.batchnorm._BatchNorm): module.eval() model.backbone.apply(set_bn_eval)这个方法在微调ImageNet预训练模型做小数据集分类时特别有价值。我之前在一批医疗影像数据上做二分类一开始直接对整个model调train()验证集精度一直在77%左右上不去改掉BN状态后同样的超参数直接跳到85%。这个坑非常值得记一笔。5. 混合精度训练和自定义更新循环参数更新最后的那一公里5.1 AMP的梯度缩放原理混合精度训练的核心是让大部分计算在FP16下进行但FP16的表示范围很小当梯度值小于约10的-4次方时会下溢成0。为了解决这个问题PyTorch的torch.cuda.amp采用“梯度缩放”在loss反传之前先乘一个缩放因子scale初始值通常是65536backward时反向传播是链式法则梯度也会被放大同样的倍数这样FP16表示范围内能保留更多有效位。scaler torch.cuda.amp.GradScaler() for inputs, labels in dataloader: optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs model(inputs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()scaler.scale(loss).backward()并不是直接把缩放后的loss传给backward后完事。它内部会记录这次backward是否产生了inf/nan。如果出现了scaler.step(optimizer)会跳过这次参数更新然后scaler.update()会把scale因子缩小通常是除以2让后续的梯度落在更安全的表示范围内。所以这几行的关系是step可能什么都不做scale因子在动态调整。5.2 unscale、clip、step、update的执行顺序与坑如果你在混合精度训练中还需要梯度裁剪不能直接写clip_grad_norm_。此时的梯度是放大后的梯度不能直接用。正确顺序是先scaler.unscale_(optimizer)把梯度反缩放到真实的数值范围再做裁剪最后scaler.step(optimizer)。scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()这里有个容易出错的点如果unscale_之后发现梯度有inf/nanstep会直接跳过更新scale因子也会在update里缩小。但如果你在step之前又手动读取了param.grad做其他操作要小心梯度可能已经被unscale_清掉了。我在实际项目里用AMP训练分割模型时就遇到过一次先clip后原来per-iteration的loss完全正常但某一轮突然出现loss跳变排查后发现是clip_grad_norm_在unscale之前执行导致裁剪作用于一个被放大了65536倍的梯度数值直接崩溃。排错过程本身很快但当时卡了半天写在这里提醒大家。5.3 手写一个简化版更新循环理解torch.optim的本质如果到现在你还想知道optimizer.step()内部最核心的动作是什么其实就是对每个参数执行一条更新公式。我们以SGD为例手写一个极简版本def sgd_step_params(params, lr, momentum0.9): for param in params: if param.grad is None: continue if momentum: # PyTorch内部维护的momentum_buffer buf param.grad.clone() param.data.add_(buf, alpha-1 * lr)真实源码里逻辑要更复杂还包含权重衰减、Nesterov、状态管理等等。但这个简化版能帮你理解所谓参数更新就是在拿到.grad后依公式修改param.data。所有优化器都是围绕“如何从grad推导出更新量”来做文章。理解了这一点自定义优化器或者特殊更新策略就变得很简单。比如你要给某层参数单独做EMA更新或想在no_grad模式下做推理与更新的混合逻辑都可以绕过torch.optim直接操作param.data。但注意手写更新时不要用param - lr * param.grad这种写法它会原地修改inplace破坏自动求导的追踪正确做法是使用.data.add_或with torch.no_grad():。另外一个心得torch.optim.lr_scheduler的step时机也直接影响参数更新节奏。PyTorch 1.x之后推荐每个epoch结束后调用scheduler.step()而torch.optim.lr_scheduler.OneCycleLR又要求在batch级别step。如果你开了学习率调度显式确认step的粒度否则可能出现“调度器认为已经过了一个epoch实际还在第一个epoch”这种降低学习率过早的问题。在PyTorch里训练模型参数更新从来不只是“调一个优化器”那么简单。它关联着自动求导的梯度流、优化器内部的状态、数值稳定性、显存策略以及模型训练与推理两种模式的状态切换。把这些环节彻底搞明白调试模型时就能少走很多弯路。上面提到的冻结BN、梯度累积的loss均分、AMP下的unscale顺序都是我实际训练里真金白银踩出来的经验希望你用不上但真要遇到了能很快反应出问题出在哪个环节。