如何无需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()
关键细节说明
- 独立优化器:彻底拆分两个网络的参数更新逻辑,确保
loss_tenc仅影响tenc,loss_unet仅影响unet,不会出现参数交叉更新的问题。 backward的inputs参数:通过指定inputs为目标网络的参数集合,限定反向传播仅计算该部分参数的梯度,PyTorch会自动释放无关的计算图节点,完全不需要retain_graph=True。- 单次前向传播:所有损失基于同一次前向的结果计算,避免重复执行高成本的模型推理,大幅提升训练效率。
- 混合精度兼容:保留原代码的
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
相关产品推荐
相关产品推荐

