PyTorch中optimizer.step()如何获取最新损失执行参数更新?
PyTorch中
optimizer.step()无需显式传入损失即可更新参数的原理 这个差异本质是PyTorch和Keras的框架设计思路不同,核心机制可以拆成三点:
- 优化器在初始化阶段就已经完成了和待优化参数的绑定
你在代码里创建优化器实例时,一定会传入类似net.parameters()的入参,比如示例配套的优化器初始化代码是optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)。这一步已经把模型所有需要训练的参数的引用地址全部存在优化器内部的参数列表里了,和Keras在model.compile阶段绑定优化器的作用完全一致,只是绑定动作的发起方从模型变成了优化器,不需要后续再重复传入模型。 - 梯度不是存在优化器或损失变量里,而是直接存在参数张量自身
PyTorch所有支持自动求导的张量都自带.grad属性,专门用来存储反向传播计算得到的梯度值。当你调用loss.backward()时,PyTorch的autograd引擎会自动沿着计算图从损失值往回追溯,把每个可训练参数对应的梯度计算出来,直接写入对应参数张量的.grad属性中。这个过程是自动求导系统和张量本身的内置能力,全程不需要优化器参与,也不需要把损失值传递给优化器。 optimizer.step()的执行逻辑根本不需要感知损失值本身
所有参数优化算法(SGD、Adam、RMSprop等)的更新逻辑,只需要三类信息就能完成:- 当前的参数值(优化器已经持有参数引用,直接读取即可)
- 当前参数对应的梯度(直接读取参数的
.grad属性即可) - 优化器自身存储的超参数(学习率、动量系数等)、历史状态(比如动量法累计的历史梯度、Adam存储的一阶/二阶矩估计值)
以最朴素的无动量SGD为例,step()的核心逻辑和下面的代码等价:
for param_group in optimizer.param_groups: lr = param_group['lr'] for p in param_group['params']: # 直接读取p.grad上的梯度更新参数,不需要接触loss p.data = p.data - lr * p.grad.data
你可以对照示例代码的执行顺序验证逻辑:先调用
optimizer.zero_grad()把所有参数上存储的上一轮旧梯度清零→前向传播算输出→计算损失→调用loss.backward()把新的梯度写入每个参数的.grad属性→调用optimizer.step()读取所有参数的梯度完成更新,整个链路完全闭环,不需要额外传入损失或标签。
内容的提问来源于stack exchange,提问作者Cranjis
相关产品推荐
相关产品推荐

