多任务学习仅在任务间存在真实语义或结构关联时有效,否则会拖慢收敛、损害精度;适用场景包括输入相同输出不同、底层感知一致高层目标不同、数据偏斜但任务耦合;pytorch中需显式分离共享主干与私有head,loss加权推荐uncertaintyweighting。

多任务学习不是“提升效率”的万能开关,而是当任务间存在真实语义或结构关联时,才能通过共享表示降低冗余、缓解小样本过拟合——用错场景反而拖慢收敛、损害单任务精度。
什么时候该用多任务学习?看任务是否真相关
强行拼凑不相关的任务(比如同时训图像分类和股票价格预测)只会让梯度冲突、loss震荡,模型更难收敛。真正适合的场景有明确共性:
- 输入相同、输出维度不同:如中文文本同时做
pos_tagging和ner,都依赖词法/句法特征 - 底层感知一致、高层目标不同:如自动驾驶中
object_detection和depth_estimation都需理解场景几何 - 数据分布偏斜但任务耦合:如推荐系统中
click(数据多)、purchase(数据少),后者可借前者特征迁移
判断标准很简单:去掉共享层后,各任务独立模型的 baseline 性能是否明显低于联合训练?如果不是,别硬上MTL。
PyTorch里怎么搭共享主干 + 多头?别直接堆nn.ModuleDict
常见错误是把所有 head 全塞进 nn.ModuleDict,结果 forward 时忘记对齐 batch 维度或 feature channel,报 RuntimeError: size mismatch。正确做法是显式分离共享与私有路径:
- 共享部分用单一 backbone(如
BertModel或resnet18.features),输出统一 shape(例如[B, D]或[B, C, H, W]) - 每个 task head 单独定义为
nn.Sequential,输入必须严格匹配共享层输出;例如 detection head 输入要是[B, 1024, 7, 7],就别接在AdaptiveAvgPool2d(1)后面 - forward 返回 dict,键名用字符串(如
"detection"),避免用数字索引导致后续 loss 计算顺序错乱
示例片段:
def forward(self, x):
shared = self.backbone(x) # shape: [B, 512, 7, 7]
return {
"cls": self.cls_head(shared.mean(dim=[2,3])), # global avg pool
"seg": self.seg_head(shared) # keep spatial
}
loss加权不能靠拍脑袋:静态权重易失效
0.7 * loss_cls + 0.3 * loss_seg 这种写法在多数场景下会迅速让强势任务主导训练。尤其当任务量纲差异大(如分类 loss≈0.5,分割 loss≈2.3),直接加权等于没平衡。
- 优先用
UncertaintyWeighting:把每个任务 loss 的 log-variance 当作可学习参数,PyTorch 实现只需几行:self.log_var_task1 = nn.Parameter(torch.zeros(())),然后 loss 改为torch.exp(-self.log_var_task1) * loss1 + self.log_var_task1 - 避免 GradNorm 在小批量或 early epoch 调用——它依赖梯度范数,batch_size
- 监控各 task loss 曲线:如果某个 loss 十几轮不降,不是权重低,很可能是 head 输入 shape 错了或 label 编码不匹配(比如 segmentation 用了
nn.CrossEntropyLoss但 mask 是 float32 0/1,该用nn.BCEWithLogitsLoss)
验证阶段必须拆开评估,不能只看总 loss 下降
训练时总 loss 下降 ≠ 所有任务变好。常见陷阱是 segmentation head 过拟合训练集边缘,val loss 涨但总 loss 因 classification 稳定而微降,误判为收敛。
- 每个 epoch 结束后,单独 run 各 task 的 validation:调用
model.eval(),再分别取 output dict 中对应 key 计算指标(如f1_score、mIoU) - 早停(early stopping)不能基于总 loss,而应选核心任务指标(如推荐场景选
purchase_auc,NLP 场景选ner_f1) - 推理时注意 head 依赖:如果 segmentation head 依赖 detection head 的 bbox crop 区域,就必须先跑 detection 再 feed 到 seg head,不能并行——这直接影响部署 latency
多任务真正的复杂点不在代码长度,而在任务边界是否清晰、label 对齐是否无歧义、以及验证逻辑是否彻底解耦。写完 model.forward 后,花三倍时间检查 data loader 输出和 eval loop,比调 learning rate 更关键。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











