从“头”说起:一个反直觉的压缩场景
假设你负责将一个训练好的英俄机器翻译模型部署到资源受限的服务器上,目标是降低推理延迟和显存占用。你首先想到的是减少模型规模,但直接缩小隐藏层维度会破坏已经学好的表示。这时你注意到,Transformer 的每一层都包含多个注意力头,每个头理论上负责关注输入的不同部分。直觉上,头越多,模型表达能力越强,但这也意味着更多的参数和计算量。
一个反直觉的发现是:在推理时移除大量注意力头,模型性能几乎不受影响,甚至有时还会提升。Paul Michel 等人在 2019 年的论文《Are Sixteen Heads Really Better than One?》中报告,即使模型用多个头训练,测试时很大比例的注意力头可以被移除而不显著影响性能,某些层甚至可以只保留一个头。Elena Voita 等人在同年的 ACL 论文《Analyzing Multi-Head Self-Attention: Specialized Heads Do the Heavy Lifting, the Rest Can Be Pruned》中进一步发现,编码器中最重要的头承担着一致且可解释的语言角色,而大部分头是冗余的。
这引出一个实际问题:既然多头注意力是 Transformer 的核心组件,为什么删掉一部分头反而能保持甚至提升性能?要回答这个问题,需要先理解注意力头在模型中到底做了什么,以及剪枝操作如何改变信息的流动。
注意力头的角色:冗余与专业化
多头注意力的设计初衷是让模型同时关注不同位置的表示子空间。每个头有自己的查询、键、值投影矩阵,独立计算注意力权重,最后将各头的输出拼接并线性变换。理论上,不同头可以学习到互补的模式,例如一个头关注句法依赖,另一个头关注共指关系。
然而,实际训练中并非所有头都发展出独特的功能。Voita 等人的工作表明,编码器中最重要的头往往具有明确且可解释的角色,例如关注相邻词、关注句法依赖中的父节点等。这些“专业化”的头对模型性能贡献最大,而其余头可能只是重复或噪声。Michel 等人的观察也支持这一点:许多头在测试时可以被移除,说明它们对最终预测的边际贡献很小。
这种冗余可能源于训练动态。多头注意力在训练初期会随机初始化,梯度下降过程中,某些头可能因为初始化位置或优化路径而占据主导,其他头则被“压制”,最终形成冗余。Michel 等人提供了初步证据,表明训练动态在多头注意力带来的增益中起作用。这意味着,冗余头并非设计缺陷,而是优化过程的自然结果。
剪枝方法:从贪心到可学习门控
识别出冗余头之后,下一步是决定如何剪枝。最简单的方法是贪心算法:逐个评估每个头的重要性,移除对验证集性能影响最小的头。Michel 等人使用了基于梯度的重要性分数,通过计算损失对每个头输出的梯度来估计其影响。这种方法在测试时直接移除头,无需重新训练,但可能不是最优的,因为头之间可能存在交互。
更系统的方法是 Voita 等人提出的基于随机门控和 L0 惩罚的剪枝方法。他们在每个头前添加一个可学习的门控变量,训练时用 L0 正则化鼓励门控值趋向 0,从而在训练过程中自动识别并剪除不重要的头。这种方法在英俄 WMT 数据集上,剪掉 48 个编码器头中的 38 个,BLEU 只下降 0.15。相比之下,贪心方法可能需要更多步骤才能达到类似效果。
两种方法的区别在于:贪心方法是在训练后基于重要性分数进行剪枝,而门控方法是在训练中学习哪些头应该被保留。后者通常能获得更好的剪枝效果,因为它考虑了头之间的协同作用,但训练成本更高。
为什么剪枝后性能不降反升
剪枝后性能提升的机制可以从两个角度解释:去噪和缓解过拟合。
首先,冗余头可能引入噪声。多头注意力的输出是各头输出的拼接,如果某些头学习到的注意力模式与任务无关,甚至与正确预测相冲突,那么这些头的输出会干扰最终表示。移除它们相当于去除了噪声源,使模型更专注于有效特征。Voita 等人的观察支持这一点:被剪掉的头往往是那些“不自信”的头,即注意力权重分布均匀、没有明确聚焦的头,它们对预测的贡献更像是随机扰动。
其次,剪枝可以视为一种正则化,缓解过拟合。模型参数越多,越容易记住训练集中的特定模式,导致泛化能力下降。剪掉部分头减少了模型容量,迫使剩余头学习更鲁棒的特征。在机器翻译等任务中,训练数据有限,过拟合是常见问题,剪枝因此可能带来泛化性能的提升。
需要注意的是,性能提升并非总是发生。Michel 等人发现,在某些层,剪掉太多头会导致性能下降,尤其是靠近输出的层。这是因为这些层的头可能承担着更关键的信息整合任务。因此,剪枝需要谨慎选择要移除的头,而不是盲目删除。
结构化剪枝与稀疏剪枝:实现层面的权衡
在部署时,剪枝方式直接影响计算图的结构和硬件利用率。注意力头剪枝有两种实现路径:结构化剪枝和稀疏剪枝。
结构化剪枝是指直接删除整个注意力头,即移除对应的查询、键、值投影矩阵的列和输出投影矩阵的行。这样,模型的计算图可以保持规整,矩阵乘法可以直接跳过被删除的维度,从而减少浮点运算量(FLOPs)。在推理时,如果使用专门的推理引擎,结构化剪枝可以带来实际的延迟降低,因为矩阵尺寸变小,内存带宽需求也降低。
稀疏剪枝则是在注意力权重矩阵中引入稀疏性,例如将某些头的注意力权重置零,但保留矩阵的维度。这种方式在数学上等效于移除头,但计算图仍然包含完整的矩阵乘法,只是其中部分元素为零。在通用硬件上,稀疏矩阵运算往往无法充分利用硬件加速,除非使用专门支持稀疏计算的库或硬件。因此,稀疏剪枝在理论上有压缩潜力,但实际加速效果有限。
下表对比了两种方式在推理部署中的关键差异:
| 维度 | 结构化剪枝 | 稀疏剪枝 |
|---|---|---|
| 计算图 | 维度减小,矩阵乘法规模变小 | 维度不变,部分元素为零 |
| 延迟收益 | 显著,尤其在 CPU 或低端 GPU 上 | 有限,除非硬件支持稀疏加速 |
| 显存占用 | 直接减少参数存储 | 需要存储完整矩阵,可能需额外索引 |
| 实现复杂度 | 简单,直接删除参数 | 复杂,需要稀疏格式和自定义算子 |
| 适用场景 | 通用部署,易于优化 | 需要稀疏库或专用硬件 |
从工程角度看,结构化剪枝更易于集成到现有推理框架中,因为它只是减小了矩阵维度,无需修改算子。稀疏剪枝则更适合研究场景,用于探索模型压缩的极限,但在实际部署中往往收益有限。
推理部署中的收益与风险
以一个实际的部署场景为例:假设你有一个 6 层 Transformer 的翻译模型,每层 8 个头,共 48 个头。通过重要性评估,你发现可以安全地剪掉 30 个头,保留 18 个。结构化剪枝后,每层的注意力计算量减少约 62.5%,因为每个头的计算量是独立的。这直接降低了推理延迟,尤其是在 batch size 较小、内存带宽受限的场景下。
然而,剪枝也带来风险。最直接的风险是精度损失。虽然许多头可以剪掉,但哪些头重要取决于任务和数据。如果剪枝决策基于验证集,而部署时的数据分布与验证集不同,那么原本不重要的头可能变得关键,导致性能下降。因此,剪枝后需要在目标数据分布上重新评估模型。
另一个风险是剪枝与微调的交互。如果剪枝后不进行微调,模型可能无法完全适应剩余头的组合;如果进行微调,则需要额外的训练成本。Voita 等人的方法在训练中学习门控,剪枝后无需微调即可保持性能,但这种方法需要从头训练或继续训练,成本较高。
此外,剪枝可能影响模型的可解释性。保留的头往往是那些有明确角色的头,但剪掉的头可能在某些输入上提供必要的备选路径。例如,在翻译罕见词或处理长句时,冗余头可能提供额外的上下文,剪掉后模型可能在这些情况下表现不佳。因此,剪枝后需要针对长尾输入进行压力测试。
剪枝流程:从评估到部署
下面是一个典型的注意力头剪枝流程,以机器翻译模型为例:
flowchart TD
A[训练好的Transformer模型] --> B[计算每个头的重要性分数]
B --> C{剪枝比例设定}
C -->|贪心| D[迭代移除最低分头]
C -->|门控| E[训练时学习门控]
D --> F[在验证集评估性能]
E --> F
F --> G{性能是否可接受?}
G -->|否| H[调整剪枝比例或微调]
H --> D
G -->|是| I[结构化剪枝生成紧凑模型]
I --> J[部署并监控线上指标]
流程从训练好的模型开始,首先计算每个头的重要性分数。分数可以基于梯度(如 Michel 等人的方法)或基于门控值(如 Voita 等人的方法)。然后根据预设的剪枝比例,移除分数最低的头。剪枝后,在验证集上评估性能,如果性能下降过多,可以调整比例或进行微调。最后,将剪枝后的模型转换为结构化剪枝版本,部署到推理环境,并持续监控线上指标,如 BLEU、延迟和显存占用。
关键步骤是重要性分数的计算。Michel 等人的方法使用梯度近似,计算损失对每个头输出的梯度,然后根据梯度范数排序。这种方法不需要重新训练,但可能忽略头之间的交互。Voita 等人的门控方法则在训练中学习每个头的保留概率,更准确但需要训练成本。
可观测信号与失败模式
部署剪枝模型后,需要监控哪些信号?首先是性能指标,如翻译质量(BLEU)是否与验证集一致。其次是延迟和吞吐量,确认剪枝确实带来了加速。此外,还应关注注意力分布的变化:如果某些层的注意力权重变得过于集中或分散,可能表明剪枝破坏了信息流动。
常见的失败模式包括:
- 分布偏移:部署数据与训练数据分布不同,导致被剪掉的头在部署数据上变得重要。例如,训练数据以新闻为主,部署时遇到大量口语化文本,某些头可能专门处理正式语体,剪掉后性能下降。
- 长尾输入:罕见词、长句或特殊句法结构可能依赖冗余头提供的额外上下文。剪枝后,模型在这些输入上可能产生错误翻译。
- 过度剪枝:在某些层,尤其是靠近输出的层,头的重要性较高,过度剪枝会导致性能急剧下降。Michel 等人的实验表明,某些层可以剪到只剩一个头,但其他层可能需要保留更多。
- 评估偏差:验证集可能无法反映真实部署场景。如果验证集与部署数据分布差异大,剪枝决策可能不准确。
针对这些失败模式,可以采取的措施包括:在目标数据上重新评估剪枝模型,保留更多头以应对长尾输入,或者使用域自适应技术调整模型。
与替代方案比较:剪枝之外的选择
注意力头剪枝并非唯一的模型压缩手段。与权重剪枝、低秩分解和知识蒸馏相比,它各有优劣。
权重剪枝是指对权重矩阵中的单个元素进行剪枝,产生稀疏矩阵。它比头剪枝更细粒度,可能保留更多信息,但稀疏矩阵在硬件上难以加速,且需要专门的稀疏库。低秩分解则将权重矩阵分解为低秩近似,减少参数数量,但可能改变模型结构,需要重新训练。知识蒸馏则是训练一个小的学生模型来模仿大模型的行为,通常能获得更好的性能,但训练成本高,且需要访问教师模型的输出。
注意力头剪枝的优势在于其结构化特性:剪枝后模型仍然是密集的,可以直接利用现有推理框架加速。它特别适合 Transformer 模型,因为多头结构天然提供了可剪枝的单元。然而,它只能移除整个头,无法对头内部的维度进行细粒度调整,因此压缩率可能不如权重剪枝。
在实际应用中,注意力头剪枝常与其他方法结合。例如,先剪掉冗余头,再对剩余权重进行低秩分解,或者用知识蒸馏进一步压缩。这种组合可以最大化压缩效果,但需要权衡实现复杂度。
尚未解决的问题
尽管注意力头剪枝在机器翻译等任务上表现出色,但仍有一些问题悬而未决。例如,如何自动确定每层的最佳剪枝比例?目前的方法通常依赖验证集,但验证集可能无法代表所有部署场景。另一个问题是,剪枝对模型鲁棒性的影响尚未被充分研究。剪掉冗余头可能使模型对对抗样本更敏感,因为冗余头可能提供了某种形式的鲁棒性。
此外,现有研究大多集中在编码器,解码器中的头剪枝研究较少。解码器的自注意力与编码器不同,可能更依赖多头机制,剪枝效果可能不同。未来工作可以探索解码器头的角色,以及剪枝对生成质量的影响。
总的来说,注意力头剪枝提供了一种有效的模型压缩手段,其性能提升的机制在于去除噪声和缓解过拟合。在部署时,结构化剪枝能带来实际的延迟和显存收益,但需要谨慎评估剪枝对长尾输入和分布偏移的影响。理解头的角色和剪枝的边界,是安全利用这一技术的关键。
资料来源
- Are Sixteen Heads Really Better than One?
- Analyzing Multi-Head Self-Attention: Specialized Heads Do the Heavy Lifting, the Rest Can Be Pruned
- Analyzing Multi-Head Self-Attention: Specialized Heads Do the Heavy Lifting, the Rest Can Be Pruned
- Losing Heads in the Lottery: Pruning Transformer Attention in Neural Machine Translation