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

如何无需retain_graph=True反向传播两个损失不同的串联网络?

串联网络分损失优化(无需retain_graph=True)的实现方案

核心思路

针对高计算成本的串联网络任务,要实现两个损失分别优化对应网络且不保留计算图,关键是:

  • 使用两个独立优化器,精准控制每个损失的更新对象;
  • 利用PyTorch反向传播的参数限定功能,避免不必要的计算图保留;
  • 仅执行一次前向传播,最大化训练效率。

完整实现代码

1. 定义独立优化器与调度器

import torch

# 为tenc和unet分别创建优化器
optimizer_tenc = torch.optim.AdamW(
    tenc.parameters(),
    lr=1e-07,
    weight_decay=1e-3,
    betas=(0.9, 0.99),
    eps=1e-07,
    fused=True,
    foreach=False
)
optimizer_unet = torch.optim.AdamW(
    unet.parameters(),
    lr=1e-06,
    weight_decay=1e-2,
    betas=(0.9, 0.99),
    eps=1e-07,
    fused=True,
    foreach=False
)

# 对应各自的学习率调度器(按需配置)
scheduler_tenc = custom_scheduler(optimizer=optimizer_tenc, warmup_steps=30, exponent=5, random=False)
scheduler_unet = custom_scheduler(optimizer=optimizer_unet, warmup_steps=30, exponent=5, random=False)

scaler = torch.cuda.amp.GradScaler()

2. 前向传播与分批次反向更新

# 单次前向传播,生成所有所需结果
with torch.cuda.amp.autocast():
    hidden_state = tenc(input)
    model_pred = unet(hidden_state)
    # 计算基础损失(无reduction,方便后续生成两个损失)
    loss = torch.nn.functional.mse_loss(model_pred, target, reduction='none')
    loss_tenc = loss.mean()
    loss_unet = (loss * mask).mean()

# 第一步:更新tenc参数,仅计算tenc的梯度
scaler.scale(loss_tenc).backward(inputs=list(tenc.parameters()))
scaler.unscale_(optimizer_tenc)
scaler.step(optimizer_tenc)
optimizer_tenc.zero_grad(set_to_none=True)  # 清空tenc梯度,避免干扰后续计算

# 第二步:更新unet参数,仅计算unet的梯度
scaler.scale(loss_unet).backward(inputs=list(unet.parameters()))
scaler.unscale_(optimizer_unet)
scaler.step(optimizer_unet)
optimizer_unet.zero_grad(set_to_none=True)  # 清空unet梯度

# 后续常规操作
scaler.update()
scheduler_tenc.step()
scheduler_unet.step()

关键细节说明

  1. 独立优化器:彻底拆分两个网络的参数更新逻辑,确保loss_tenc仅影响tenc,loss_unet仅影响unet,不会出现参数交叉更新的问题。
  2. backward的inputs参数:通过指定inputs为目标网络的参数集合,限定反向传播仅计算该部分参数的梯度,PyTorch会自动释放无关的计算图节点,完全不需要retain_graph=True。
  3. 单次前向传播:所有损失基于同一次前向的结果计算,避免重复执行高成本的模型推理,大幅提升训练效率。
  4. 混合精度兼容:保留原代码的GradScaler与autocast逻辑,确保训练的数值稳定性与硬件利用率。

备选方案(手动梯度计算)

如果需要对梯度做额外处理(如裁剪、加权),可以使用torch.autograd.grad手动计算梯度:

with torch.cuda.amp.autocast():
    hidden_state = tenc(input)
    model_pred = unet(hidden_state)
    loss = torch.nn.functional.mse_loss(model_pred, target, reduction='none')
    loss_tenc = loss.mean()
    loss_unet = (loss * mask).mean()

# 计算并更新tenc梯度
tenc_grads = torch.autograd.grad(scaler.scale(loss_tenc), tenc.parameters())
for param, grad in zip(tenc.parameters(), tenc_grads):
    param.grad = grad
scaler.unscale_(optimizer_tenc)
scaler.step(optimizer_tenc)
optimizer_tenc.zero_grad(set_to_none=True)

# 计算并更新unet梯度
unet_grads = torch.autograd.grad(scaler.scale(loss_unet), unet.parameters())
for param, grad in zip(unet.parameters(), unet_grads):
    param.grad = grad
scaler.unscale_(optimizer_unet)
scaler.step(optimizer_unet)
optimizer_unet.zero_grad(set_to_none=True)

scaler.update()
scheduler_tenc.step()
scheduler_unet.step()

内容的提问来源于stack exchange,提问作者Anonymous

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 00:54:52