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(低资源) |
| 每任务单独 BERT | 61.50% | 72.33% | 34.40% |
| 100 任务 MTL | 73.29% | 69.30% | 54.39% |
| 500 任务 MTL | 72.54% | 67.36% | 52.80% |
| 1000 任务 MTL | 74.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 已经够」这一层,没有更细的过拟合诊断。作者自己也说下一步要看异质任务和对抗任务。