You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 22:36:23