PyTorch多任务损失共享主干模型的反向传播方法咨询
多任务模型损失反向传播方案优化
问题概述
模型包含共享主干网络,两个分支分别处理二分类任务(BCELoss)和回归任务(MSELoss)。当前采用两次反向传播(第一次保留计算图)的方式,但模型学习效果不佳,需确认实现正确性并获取最优方案。
模型代码
import torch import torch.nn as nn class First_branch(nn.Module) : def __init__(self, input_size, num_heads=3): super(First_branch, self).__init__() self.fc1 = nn.Linear(input_size*num_heads, 128) self.fc2 = nn.Linear(128, 64) def forward(self, x): x = torch.relu(self.fc1(x)) x = self.fc2(x) return x class Final_branch_1(nn.Module): def __init__(self, input_size, num_heads=3): super(Final_branch_1, self).__init__() self.fc1 = nn.Linear(64, 128) self.fc2 = nn.Linear(128, 64) self.fc3 = nn.Linear(64, 1) def forward(self, x): x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) x = torch.sigmoid(self.fc3(x)) # Binary classification return x class Final_branch_2(nn.Module): def __init__(self, input_size, num_heads=3): super(Final_branch_2, self).__init__() self.fc1 = nn.Linear(64, 128) self.fc2 = nn.Linear(128, 64) self.fc3 = nn.Linear(64, 1) def forward(self, x): x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) x = self.fc3(x) #Regression task return x class Model(nn.Module): def __init__(self, in_feats_dict, hidden_feats_dict, out_feats_dict, edge_feats_dict, rel_names): super().__init__() self.conv = ConvModel(in_feats_dict, hidden_feats_dict, out_feats_dict, edge_feats_dict, rel_names, num_heads=3) self.embed = First_branch(sum(out_feats_dict.values()), num_heads=3) self.pred_1 = Final_branch_1(64, num_heads=3) self.pred_2 = Final_branch_2(64, num_heads=3) def forward(self, g, node_features, edge_features): conv_output = self.conv(g, node_features, edge_features) to_concat = [conv_output[key] for key in conv_output.keys()] # Aggregate the results following each latent feature aggregated_features = [torch.mean(i, dim=0) for i in to_concat] concatenated_features = torch.cat(aggregated_features) embedded_features = self.embed(concatenated_features) prediction_1 = self.pred_1(embedded_features) prediction_2 = self.pred_2(embedded_features) return prediction_1, prediction_2
核心结论
首先纠正一个误解:BCELoss和MSELoss完全可以直接求和,两者都是标量损失,不存在类型兼容问题。你当前的两次反向传播方式是正确的,但并非最优,且模型效果差的根源大概率不在反向传播方式,而是损失平衡、数据预处理或训练细节问题。
最优反向传播实现
方式1:损失加权求和后单次反向传播(推荐)
这是多任务学习的标准做法,通过权重参数平衡两个任务的梯度贡献,避免单一任务损失主导训练。代码示例:
from torch.nn import MSELoss # 初始化损失函数 bce_criterion = nn.BCELoss() mse_criterion = MSELoss() # 前向传播 pred1, pred2 = model(g, node_features, edge_features) # 计算损失并加权求和(权重需根据任务调整) alpha = 0.6 # 分类任务权重 beta = 0.4 # 回归任务权重 loss_cls = bce_criterion(pred1, label_cls.float()) # 确保标签为float类型 loss_reg = mse_criterion(pred2, label_reg) total_loss = alpha * loss_cls + beta * loss_reg # 反向传播与优化 optimizer.zero_grad() total_loss.backward() optimizer.step()
方式2:两次反向传播(当前实现)
这种方式本质上与求和后反向传播等价(梯度具有可加性),但写法繁琐且易出错。需注意:必须在两次反向传播前统一执行optimizer.zero_grad(),且第一次反向传播需保留计算图:
optimizer.zero_grad() loss_cls = bce_criterion(pred1, label_cls.float()) loss_reg = mse_criterion(pred2, label_reg) loss_cls.backward(retain_graph=True) loss_reg.backward() optimizer.step()
此方式无性能优势,不推荐使用。
模型效果差的关键优化方向
损失量级平衡:
- BCELoss的输出范围通常在00.7左右,而MSELoss的范围取决于回归标签的数值范围(如标签是0100,MSELoss可能达到数十甚至上百)。若不做处理,回归任务损失会主导梯度更新。
- 解决方案:将回归标签归一化到0~1范围,或调整权重参数alpha/beta,让两个损失的量级接近。
分类任务细节检查:
- BCELoss要求目标标签为
float类型(如0.0/1.0),若标签是int类型需转换:label_cls = label_cls.float()。 - 确保分类分支的输出经过
sigmoid(你的代码已实现),输出范围在0~1之间,符合BCELoss的输入要求。
- BCELoss要求目标标签为
训练细节优化:
- 多任务学习通常需要更小的学习率,可尝试从
1e-4开始调整。 - 检查训练数据:分类任务是否存在类别不平衡?回归任务是否有异常值?需针对性做数据预处理(如过采样、异常值剔除)。
- 验证共享主干的特征是否适配两个任务:可尝试在分支前添加独立的特征转换层,或调整主干网络的复杂度。
- 多任务学习通常需要更小的学习率,可尝试从
内容的提问来源于stack exchange,提问作者TheauLep
相关产品推荐
相关产品推荐

