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

基于TensorFlow实现多任务文本模型推理时动态切换任务头

解决方案:构建动态任务路由的TensorFlow推理模型

完全可以实现你的需求——将200个共享权重的文本模型整合成一个单一TensorFlow模型,通过输入(文本,任务索引i)返回对应任务的预测结果,既节省存储体积,又能仅计算目标任务头以减少运算量,且适配Google AI Platform的部署要求。

以下是具体实现步骤和代码示例:

1. 复用共享编码器权重

首先提取并加载训练好的共享文本编码器(比如自定义的文本特征提取层,或是BERT这类预训练模型的主体部分)。这部分权重所有任务共用,只需存储一次。

2. 整合独立训练的任务头

将200个单独训练好的任务头(分类层、回归层等)整理为可索引的集合,确保每个任务头的输入维度与编码器的输出维度完全匹配(迁移学习训练时应该已经满足此条件)。

3. 构建动态任务路由模型

创建自定义TensorFlow模型,接收文本输入和任务索引i两个输入,通过动态路由仅调用目标任务头计算结果:

  • 用共享编码器处理文本,生成通用特征向量
  • 根据任务索引i,精准选择对应的任务头进行预测计算
  • 输出该任务的预测结果

代码示例(TensorFlow 2.x)

import tensorflow as tf
from tensorflow.keras.layers import Input
from tensorflow.keras.models import Model

# 加载预训练好的共享编码器(替换为你的编码器路径)
shared_encoder = tf.keras.models.load_model('path/to/trained_shared_encoder')

# 加载所有独立训练的任务头,存储为字典(key为任务索引i)
task_heads = {}
for task_idx in range(200):
    task_heads[task_idx] = tf.keras.models.load_model(f'path/to/task_head_{task_idx}')

# 构建多任务推理模型
def build_multi_task_inference_model(shared_encoder, task_heads):
    # 文本输入:根据你的编码器需求定义,示例为BERT的输入格式
    input_ids = Input(shape=(None,), dtype=tf.int32, name='input_ids')
    attention_mask = Input(shape=(None,), dtype=tf.int32, name='attention_mask')
    # 任务索引输入
    task_idx_input = Input(shape=(), dtype=tf.int32, name='task_idx')

    # 共享编码器生成特征向量
    encoder_outputs = shared_encoder([input_ids, attention_mask])
    # 取编码器的池化输出(根据你的编码器结构调整)
    pooled_features = encoder_outputs['pooled_output']

    # 将所有任务头的权重堆叠为可索引张量(以单Dense层任务头为例)
    task_kernels = tf.stack([head.trainable_weights[0] for head in task_heads.values()])
    task_biases = tf.stack([head.trainable_weights[1] for head in task_heads.values()])

    # 根据任务索引选择对应权重
    selected_kernel = tf.gather(task_kernels, task_idx_input)
    selected_bias = tf.gather(task_biases, task_idx_input)

    # 计算目标任务的预测结果
    predictions = tf.matmul(pooled_features, selected_kernel) + selected_bias
    # 若为分类任务,可添加softmax:predictions = tf.nn.softmax(predictions)

    # 定义完整模型
    return Model(
        inputs=[input_ids, attention_mask, task_idx_input],
        outputs=predictions,
        name='multi_task_text_inference_model'
    )

# 生成并保存模型(用于Google AI Platform部署)
multi_task_model = build_multi_task_inference_model(shared_encoder, task_heads)
multi_task_model.save('path/to/final_multi_task_model')

关键注意事项

  • 高效路由:避免用tf.cond逐个判断任务索引,改用tf.gather索引堆叠的任务头权重,更适配TensorFlow图模式,提升推理效率。
  • 部署适配:在Google AI Platform部署时,需明确输入签名,示例输入格式如下:
    {
        "input_ids": [101, 2023, 3014, ..., 102],
        "attention_mask": [1, 1, 1, ..., 1],
        "task_idx": 12
    }
    
  • 权重冻结:推理阶段可将所有权重设置为不可训练,避免部署时意外更新:
    multi_task_model.trainable = False
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 09:27:28