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

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()

此方式无性能优势,不推荐使用。


模型效果差的关键优化方向

  1. 损失量级平衡:

    • BCELoss的输出范围通常在00.7左右,而MSELoss的范围取决于回归标签的数值范围(如标签是0100,MSELoss可能达到数十甚至上百)。若不做处理,回归任务损失会主导梯度更新。
    • 解决方案:将回归标签归一化到0~1范围,或调整权重参数alpha/beta,让两个损失的量级接近。
  2. 分类任务细节检查:

    • BCELoss要求目标标签为float类型(如0.0/1.0),若标签是int类型需转换:label_cls = label_cls.float()。
    • 确保分类分支的输出经过sigmoid(你的代码已实现),输出范围在0~1之间,符合BCELoss的输入要求。
  3. 训练细节优化:

    • 多任务学习通常需要更小的学习率,可尝试从1e-4开始调整。
    • 检查训练数据:分类任务是否存在类别不平衡?回归任务是否有异常值?需针对性做数据预处理(如过采样、异常值剔除)。
    • 验证共享主干的特征是否适配两个任务:可尝试在分支前添加独立的特征转换层,或调整主干网络的复杂度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 14:24:55