单模型扛下1000个电商分类,低资源任务准确率从34%提到57%

Exceeding the Limits of Visual-Linguistic Multi-Task Learning

Cameron R. Wolfe, Keld T. Lundgaard

cs.AI, cs.CL, cs.CV, cs.LG

2021-07-28

Rice与Salesforce用共享多模态BERT同时训练最多1000个电商层级分类,相对单任务基线,平均准确率从61.5%升到75.0%,低资源任务从34.4%升到57.1%。

这篇在解决什么

当时视觉-语言多任务学习通常一次只绑十几个任务,BERT 的容量被当成「够大」,但没人把它推到上百、上千个分类头。工业里更常见的场景是:同一家 SaaS 对着几百个客户站点,输入形态差不多(标题、描述、商品图),标签体系各不相同,小站点样本少到单独微调会塌。Rice 大学 Cameron Wolfe 在 Salesforce Einstein 实习期间,用真实电商多租户数据把这个问题做成「大规模 MTL」:单个共享模型同时解至少 100 个、最多 1000 个分类任务。

数据不能公开。100 任务集来自 50 个站点的商品 type 和 category,超过 25 万件商品,任务之间因为服饰层级相近而相关。更大集合只是多加站点,评估仍只报原来那 100 个任务,避免换测试集。

方法

骨架是 Kiela 等人的单流多模态 BERT base,用预训练权重初始化。文本按 BERT 分词;图片先过冻结的 EfficientNet-B4,再线性投到隐空间,所有图 embedding 和文本 token 拼成一条序列,图和文用不同 token type,图没有顺序所以共用同一个 position embedding。BERT 参数全任务共享,每个任务一个分类头,前向必须指定任务。

100 任务训 15 个 epoch,batch 64,一整个 batch 只来自一个任务,AdamW,单卡 V100。收敛靠学习率预热:前 4 个 epoch 从 \(10^{-5}\) 升到 \(10^{-4}\),第 8、12 个 epoch 再各降 10 倍。固定小学习率也能收敛但慢;先冻骨干再解冻略好,预热最好。

任务采样概率 \(P(T)\propto NT^{\alpha}\)。\(\alpha=1\) 按数据量采样,\(\alpha=0\) 均匀采样。指数把 \(\alpha\) 从 1.0 收到 0.1,前期吃大任务加快收敛,后期抬小任务。连续 10 步抽同一任务会在 100 任务集上直接发散,所以每一步都换任务。

全连接分类头在 100 任务上就要 5800 万任务专属参数。改成把 BERT 输出投到低维 \(dt\) 再做自注意力,\(dt=64\) 大约少 10 倍参数,准确率几乎不掉。高资源任务喜欢更大的头,低资源任务头太大就过拟合。DyPA 用标签量当复杂度代理,按四分位分配 \(dt=128/256/512/1024\),高资源头变大、低资源头保持小,任务专属参数仍比全连接少约 3.5 倍。

结果

所有数字都在原来 100 个任务的 80/20 测试集上。

方法平均准确率T10(高资源)B10(低资源)
每任务单独 BERT61.50%72.33%34.40%
100 任务 MTL73.29%69.30%54.39%
500 任务 MTL72.54%67.36%52.80%
1000 任务 MTL74.98%67.02%57.08%

单独训练在高资源任务上高 3 个点左右,那些任务每门超过 1 万条标签,本来就能自己撑住。低资源端单独训练低于 35%,1000 任务模型把 B10 抬了 22.7 个点。1000 个独立 BERT base 大约 1100 亿参数,这个共享模型大约 2.5 亿。

把 MTL 骨干当预训练,迁到 405,840 条、2,196 类的跨站商品本体分类:BERT base 90.27%,100 任务 MTL 90.77%,1000 任务 91.12%。换 BERT large 当共享骨干,100 任务平均准确率 73.39%,对 T10 和 B10 都没超过 base。再塞进一个比任何原任务大约一个数量级的本体分类任务,T10 反而升约 4 个点,采样和 DyPA 没有被大任务挤垮。10、25、50、75 任务的小规模消融里,\(\alpha\) 衰减加 DyPA 也总是最好,增益大约 4.5–9 个点。

为什么重要

给多租户分类一个能落地的配方:共享 BERT、每步换任务、\(\alpha\) 指数衰减、按数据量给任务头分配容量。小站点借到大站点的层级归纳,大站点只让出大约 3 个 T10 点。2021 年的设定放在今天看仍清楚:任务高度相关、输入同构,这是 MTL 最顺的坡,不是任意 1000 个互不相关任务都能这么压。

局限与存疑

数据私有,外人无法复现或检查标签噪声。任务几乎都是电商商品层级,相关度高,外推到 GLUE 那种异质任务组没有证据。图片骨干冻结,多图无序。评估是每任务准确率再平均,没有校准、长尾 F1 或公平性。1000 任务模型的测试集仍是原来那 100 门,多出来的 900 门只当训练信号。BERT large 没带来好处,容量故事停在「base 已经够」这一层,没有更细的过拟合诊断。作者自己也说下一步要看异质任务和对抗任务。

术语

原文与代码

社区讨论

相关论文

全部论文解读