多任务“嵌套”神经网络实现方法相关技术问询
多任务神经网络复现问题
我正在尝试复现一篇论文中使用的多任务神经网络,但作者未提供对应部分的代码,我目前不清楚该如何编写该多任务网络的实现代码。
该网络架构如下:
为简化理解,该网络架构可概括为(为便于演示,我将原文中一对独立embedding的复杂运算替换为了拼接操作):
论文作者将单任务损失和配对任务损失相加,使用总损失在每个batch中优化encoder、MLP-1、MLP-2三个网络的参数,但我不清楚如何在单个batch中组合不同类型的数据,输入到共享初始encoder的两个不同网络中。我搜索过类似结构的其他网络但未找到相关资料,恳请各位给出相关建议,谢谢!
实现方案(基于PyTorch框架示例)
核心逻辑
- 数据组织:每个batch同时包含两类数据,分别是用于计算单任务损失的
N条单样本数据,以及用于计算配对损失的M组配对样本数据,每组配对样本由两个有同/异标签关联的单样本组成。 - 前向传播:将所有单样本(包括独立的单任务样本、配对任务拆分出的两个样本)统一输入共享encoder做特征提取,之后拆分特征分别输入两个MLP头计算对应损失,总损失为两个损失的加权和。
- 反向传播:直接对总损失调用反向传播方法即可,共享encoder的梯度会自动累加两个任务路径的梯度,不需要额外处理。
示例代码
import torch import torch.nn as nn # 共享编码器 class Encoder(nn.Module): def __init__(self, input_dim, hidden_dim): super().__init__() self.layers = nn.Sequential( nn.Linear(input_dim, hidden_dim * 2), nn.ReLU(), nn.Linear(hidden_dim * 2, hidden_dim) ) def forward(self, x): return self.layers(x) # 单任务头MLP-1 class SingleTaskHead(nn.Module): def __init__(self, hidden_dim, out_dim): super().__init__() self.layer = nn.Linear(hidden_dim, out_dim) def forward(self, x): return self.layer(x) # 配对任务头MLP-2 class PairTaskHead(nn.Module): def __init__(self, hidden_dim, out_dim): super().__init__() # 两个特征拼接所以输入维度是2*hidden_dim self.layers = nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, out_dim) ) def forward(self, x1, x2): concat_feat = torch.cat([x1, x2], dim=-1) return self.layers(concat_feat) # 单batch训练逻辑 def train_step(single_x, single_y, pair_x1, pair_x2, pair_y, encoder, mlp1, mlp2, optimizer): optimizer.zero_grad() # 合并所有样本一次性过编码器,提升计算效率 all_input = torch.cat([single_x, pair_x1, pair_x2], dim=0) all_feat = encoder(all_input) # 拆分特征 single_feat = all_feat[:len(single_x)] pair_feat1 = all_feat[len(single_x): len(single_x) + len(pair_x1)] pair_feat2 = all_feat[len(single_x) + len(pair_x1): ] # 计算两个损失 loss_single = nn.CrossEntropyLoss()(mlp1(single_feat), single_y) loss_pair = nn.CrossEntropyLoss()(mlp2(pair_feat1, pair_feat2), pair_y) # 总损失反向传播更新参数 total_loss = loss_single + loss_pair total_loss.backward() optimizer.step() return total_loss.item()
内容的提问来源于stack exchange,提问作者user48867
相关产品推荐
相关产品推荐

