首页
学习
活动
专区
圈层
工具
发布

一个编码器加七个任务头:我们训练统一安全分类器的经验总结

过去几个月,我们把原本各自独立的七个序列分类模型,整合进了一个多任务头的统一模型,也就是我们现在的旗舰模型既然模型权重已经公开,我就来分享一下哪些做法有效,以及有哪些结果出乎我们的意料

先说模型结构:我们用了一个共享的 mmBERT-small 编码器,上面接七个任务头,分别是二进制注入检测(BCE)文档类别(7类)工具类型(14类)工具操作(6类)工具数据流标签(3个多标签 BCE)意图路由(5类)和威胁类型(7类)

需要特别小心的地方在于:训练数据里每行只带有部分任务的标签,所以那些没标注的任务,在计算损失时会被完全屏蔽掉我们专门写了一个自检测试,确认被屏蔽任务的梯度必须严格为零,结果还真揪出了两个隐蔽的 bug如果你也在做类似的损失屏蔽,建议也加上这个测试另外,大约 5000 条合成加真实的多任务数据帮七个任务头更好地协同训练,而测试集则全部用的真实数据

各任务头在预留测试集上的表现如下:注入检测 F1 0.962,文档分类 0.980,工具类型 0.957,工具操作 0.945,工具标签 0.958,意图路由 0.916,威胁类型 0.952

量化方面,统一模型和单独的单任务模型都提供了量化后的 edge 版本(ONNX INT8 加 INT4 嵌入,原始约 96MB),代码库里附有量化前后的一致性评测,表现最差的任务头相比 FP32 只掉了 0.012

那比起七个独立模型,统一方案到底值不值?两种版本我们都放出来了,大家可以自行对比单任务模型在多数任务上分数略高一点,但统一模型只需要跑一次编码器,而不是最多七次

我们目前的短板是意图路由,只有 0.916主要是意图类别之间语义有重叠,比如“写一段分析我数据的代码”到底算编程还是分析,本身就模糊我怀疑这种歧义是数据里真实存在的如果你有除了重新标注之外的好办法,欢迎交流

  • 发表于:
  • 原文链接https://page.om.qq.com/page/OPZggeNXKmXIFwgNBFEsDaBg0
  • 腾讯「腾讯云开发者社区」是腾讯内容开放平台帐号(企鹅号)传播渠道之一,根据《腾讯内容开放平台服务协议》转载发布内容。
  • 如有侵权,请联系 cloudcommunity@tencent.com 删除。

相关快讯

领券