Meta 提出字节级蒸馏方法,预测能力上限高于 Token 蒸馏
Meta 与华盛顿大学提出字节级蒸馏,让学生从字节概率分布学起,按缩放定律预测的能力上限比 Token 蒸馏高 4 个百分点。
Meta FAIR 与华盛顿大学提出字节级蒸馏:把学习单位从 Token 换成只有 256 种取值的字节,避开十几万候选的 Token 词表。以 Llama 3-8B 为教师的实验按缩放定律预测,其下游平均准确率上限比传统 Token 蒸馏高 4 个百分点。
正文摘录
最近,来自 Meta FAIR 和华盛顿大学的论文《突破 Token 天花板:蒸馏出更小、更强的字节模型》,提出了一个很有趣的想法—— 蒸馏的时候,能不能把学习单位从词元(Token)换成字节,让学生模型直接从字节(byte)级别的概率分布学起? 团队给每个 Token 加上结束标记,把教师的 Token 概率分布完整转换到字节级别。在以 Llama 3-8B 为教师的蒸馏实验中,团队根据缩放定律预测: 随着训练计算量增加,字节级蒸馏将反超 Token 级蒸馏,下游平均准确率上限高出 4 个百分点。 传统蒸馏让学生学习教师在庞大 Token 词表上的概率分布。相比之下,字节只有 256 种基础取值,单个位置的预测空间更小,基础取值也不随分词器改变。 其实到今天,蒸馏已经不是什么陌生技术了。我们都知道,大的模型能力更强,但把它部署到实际应用中,显存、计算和响应速度都是成本。 和普通监督训练只告诉模型“下一个正确 Token 是什么”不同,蒸馏还会让学生学习教师对各个候选 Token 的概率判断。 但问题也出在这里。Token 的候选空间实在太大。以 Llama 3-8B 为例,它的词表包含 128256 个 Token,这意味着每一个预测位置,都对应十几万个候选概率。 如果做离线蒸馏,把完整概率分布全部保存下来,存储成本会非常高。因此实际操作中,通常只能保留概率最高的一部分,也就是 top-k 截断。 一个字节由 8 个比特组成,共有 256 种可能的取值。即使加上少量特殊符号,每个位置需要保存的概率也不过两百多个,完整分布自然更容易保留下来。 不过,教师预测的是整个 Token,学生预测的是单个字节。要让两者对上,不能只把文本拆开,还得把教师的概率分布一起转换过去。 具体的,团队提出了两种方案,分别叫作 Marginalize-It 和 End-Of-Token。