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

多任务“嵌套”神经网络实现方法相关技术问询

多任务神经网络复现问题

我正在尝试复现一篇论文中使用的多任务神经网络,但作者未提供对应部分的代码,我目前不清楚该如何编写该多任务网络的实现代码。
该网络架构如下:
网络架构
为简化理解,该网络架构可概括为(为便于演示,我将原文中一对独立embedding的复杂运算替换为了拼接操作):
简化版本架构
论文作者将单任务损失和配对任务损失相加,使用总损失在每个batch中优化encoder、MLP-1、MLP-2三个网络的参数,但我不清楚如何在单个batch中组合不同类型的数据,输入到共享初始encoder的两个不同网络中。我搜索过类似结构的其他网络但未找到相关资料,恳请各位给出相关建议,谢谢!


实现方案(基于PyTorch框架示例)

核心逻辑

  1. 数据组织:每个batch同时包含两类数据,分别是用于计算单任务损失的N条单样本数据,以及用于计算配对损失的M组配对样本数据,每组配对样本由两个有同/异标签关联的单样本组成。
  2. 前向传播:将所有单样本(包括独立的单任务样本、配对任务拆分出的两个样本)统一输入共享encoder做特征提取,之后拆分特征分别输入两个MLP头计算对应损失,总损失为两个损失的加权和。
  3. 反向传播:直接对总损失调用反向传播方法即可,共享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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 21:45:04