多任务模型需共享底层特征提取层,用函数式api构建单个model并统一training模式;损失须显式加权平衡,label格式与输出名严格对齐;predict返回列表而非字典,需封装wrapper转字典。

多任务模型必须共享底层特征提取层
TensorFlow 中实现多任务学习,核心不是堆叠多个模型,而是让不同任务共用一部分网络参数。如果每个任务都从头训练独立模型,就失去了多任务学习“知识迁移”和“正则化”的意义。典型结构是:一个共享的 base_model(比如 CNN 或 Transformer 编码器),后面接多个任务专属的 head(如分类层、回归层)。
常见错误是把多个 Model 实例拼在一起,结果梯度无法反向传播到共享层——TensorFlow 会报 ValueError: No gradients provided for any variable。正确做法是用函数式 API 构建单个 Model,所有输出分支都从同一中间层引出。
- 共享层输出必须是张量(
tf.Tensor),不能是另一个Model的输出(除非用tf.keras.layers.Lambda包装) - 多个
output在调用tf.keras.Model(inputs=..., outputs=[...])时需传入 Python 列表,顺序要和后续loss、loss_weights对齐 - 避免在共享层里使用
Dropout或BatchNormalization的训练模式不一致问题:确保整个模型统一用training=True或training=False调用
如何设置多输出损失与权重平衡
TensorFlow 默认不支持“一个模型多个 loss 自动加权”,必须显式指定 loss 字典或列表,以及可选的 loss_weights。否则会报错 ValueError: The model cannot be compiled because it has multiple outputs。
损失权重不是超参调优的“装饰项”,它直接影响梯度幅值。例如一个回归任务 loss 值常在 100+,而分类任务交叉熵在 0.5 左右,若不加权,回归任务梯度会主导更新,导致分类 head 几乎不学习。
- 用字典方式指定 loss:
loss={'task_a': 'sparse_categorical_crossentropy', 'task_b': 'mse'} - 对应权重写成字典:
loss_weights={'task_a': 1.0, 'task_b': 0.3},或按输出顺序用列表:loss_weights=[1.0, 0.3] - 自定义 loss 函数必须返回标量
tf.Tensor,且内部不要用tf.print或print,否则训练会变慢甚至卡死
训练时 label 数据格式必须对齐输出名
如果你的模型输出定义为 outputs=[output_a, output_b],且给它们起了名字:output_a = Dense(..., name='sentiment')(x),那么训练时传入的 y 必须是字典:{'sentiment': y_sentiment, 'regression_target': y_reg}。否则 TensorFlow 会找不到匹配 key,报错 KeyError: 'sentiment' 或静默忽略某任务。
更隐蔽的问题是 shape 不匹配:比如分类任务期望 (None,) 的整数标签,但你传了 (None, 1) 的二维数组,会触发 InvalidArgumentError: logits and labels must have the same shape。
- 检查每个输出的
name属性:model.output_names返回列表,和你传入y的 key 或顺序严格对应 - 用
np.squeeze()清理冗余维度,尤其从 pandas 或 h5 文件读数据后容易多一维 - 如果某个任务 batch 内无有效样本(如 mask 掉全部),loss 可能为 NaN,建议在自定义 loss 中加
tf.debugging.check_numerics
预测阶段要注意输出顺序与命名一致性
model.predict() 返回的是 Python 列表,不是字典,顺序严格对应构建模型时 outputs=[...] 的顺序。即使你在定义层时写了 name='xxx',预测结果也不会自动变成字典——这点和训练时的 label 输入逻辑不同,容易混淆。
实际部署中,如果下游服务依赖字段名(比如 JSON 接口返回 {"score": ..., "label": ...}),靠顺序维护极易出错。建议封装一层 predict wrapper。
- 预测后手动转字典:
preds = dict(zip(model.output_names, model.predict(x_test))) - 避免在 predict 前调用
model.trainable = False后忘记恢复,否则 BatchNorm 层统计量不会更新,影响推理稳定性 - 导出 SavedModel 时,
signatures需显式指定输入输出映射,否则加载后无法按名取输出
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











