PyTorch中如何用W_p^p次梯度通过SGD更新模型参数
核心原理
PyTorch中optimizer.step()的执行逻辑不依赖自动微分框架:它只会遍历所有绑定的可训练参数,读取每个参数挂载的.grad属性存储的梯度值,按照优化器对应的更新规则(比如SGD的梯度下降、带动量/权重衰减的更新逻辑)修改参数值。常规流程里调用loss.backward(),本质上只是自动完成了「计算损失对各参数的梯度、把梯度写入对应参数的.grad属性」这一步,你完全可以手动把自定义算法算出的梯度写入.grad,效果和自动求导得到的梯度没有任何区别。
具体实现流程
- 梯度清零:和常规训练一致,先调用
optimizer.zero_grad()清空上一轮迭代残留的梯度值,避免梯度累加错误。 - 前向传播:正常执行模型前向计算,得到当前batch输入对应的模型输出
y_pred,这一步不需要为Wasserstein距离的计算过程维护自动微分计算图,可以关掉相关跟踪减少内存开销。 - 梯度计算:运行你复现的算法1,得到带熵正则的$W_p^p$损失对模型输出
y_pred的次梯度。如果你不想手动逐层计算链式法则反传得到所有模型参数的梯度,可以直接调用PyTorch内置的torch.autograd.grad接口,传入模型输出、模型参数、以及算法1算出的对输出的梯度,就能自动得到所有参数对应的梯度,不需要手动实现反向传播逻辑。 - 梯度挂载:遍历所有模型可训练参数,把计算得到的对应形状的梯度赋值给参数的
.grad属性,注意保证梯度和参数的数据类型、所在设备(CPU/GPU)完全一致。 - 参数更新:直接调用
optimizer.step(),优化器会自动读取挂载好的梯度完成参数更新,和常规训练流程完全兼容。
代码示例
# 已复现的论文算法1实现 # 输入:模型预测输出y_pred、样本真实标签y、Wasserstein距离阶数p、熵正则系数epsilon # 输出:损失对y_pred的次梯度grad_wrt_output、标量损失值loss(仅用于日志记录) def wasserstein_subgradient(y_pred, y, p=2, epsilon=1e-3): # 此处填充算法1的具体计算逻辑 # ... 求解最优运输计划、计算次梯度 ... return grad_wrt_output, loss # 单步训练流程 optimizer.zero_grad() # 前向传播得到模型预测 y_pred = model(train_batch_x) # 运行算法得到对模型输出的梯度 grad_output, batch_loss = wasserstein_subgradient(y_pred, train_batch_y) # 借助自动微分从输出梯度反传得到所有模型参数的梯度 param_grads = torch.autograd.grad( outputs=y_pred, inputs=model.parameters(), grad_outputs=grad_output ) # 将梯度挂载到对应参数上 for param, g in zip(model.parameters(), param_grads): param.grad = g # 此处可插入梯度裁剪等自定义操作,例如 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # 执行参数更新 optimizer.step() # 记录训练损失 print(f"Current batch loss: {batch_loss.item():.4f}")
常见注意点
- 不需要再调用
loss.backward(),否则会重复计算梯度,覆盖或者累加你手动挂载的梯度,导致更新错误。 - 如果你的算法1实现已经直接算出了损失对每个模型参数的梯度,可以跳过
torch.autograd.grad的步骤,直接把算好的梯度赋值给对应参数的.grad即可。 - 梯度裁剪、权重衰减、学习率调度等所有常规训练配套操作,都可以在梯度挂载完成后、调用
optimizer.step()之前正常使用,逻辑和自动求导场景完全一致。 - 如果需要记录训练过程的损失值,可以在算法1里同步计算不带反传的标量损失,用来打日志或者做指标评估,不影响参数更新流程。
内容的提问来源于stack exchange,提问作者aria
相关产品推荐
相关产品推荐

