Prune Once for All: Sparse Pre-Trained Language Models 解读
一、论文基本信息论文题目Prune Once for All: Sparse Pre-Trained Language Models作者Ofir Zafrir、Ariel Larey、Guy Boudoukh、Haihao Shen、Moshe Wasserblat单位主要来自Intel Labs / Intel Corporation发表形式arXiv 2021也有 NeurIPS 2021 Efficient Natural Language and Speech Processing workshop 版本。论文提出的方法简称为Prune OFA。需要注意这里的 “Once for All” 和 CNN 领域的 OFA 超网络不是同一类方法这里的意思是预训练模型只剪一次然后这个稀疏预训练模型可以迁移到多个下游任务。(arXiv)官方实现方面论文说明作者发布了模型压缩研究库搜索结果中对应仓库是Model Compression Research Package用于支持 pruning、distillation、quantization 等压缩实验。(GitHub)二、论文要解决的问题在这篇论文之前BERT 剪枝大致有两种常见路线。第一种是下游任务剪枝。也就是对每个任务单独 fine-tune然后在这个任务上剪枝。例如 SQuAD 剪一次MNLI 剪一次QQP 再剪一次。这种方法通常效果不错但问题是每个任务都要重新调剪枝超参数训练成本高工程部署也麻烦。第二种是预训练阶段剪枝。也就是先得到一个稀疏 BERT再把它迁移到不同下游任务。这个方向更有吸引力因为一旦得到一个通用稀疏预训练模型后面很多任务都可以直接使用。论文明确引用 Gordon 等人的结论BERT 在预训练阶段剪枝和在 fine-tuning 阶段剪枝相比最终任务精度差别不大这为“只剪一次、迁移多任务”提供了动机。所以这篇论文要解决的问题是能不能在预训练阶段把 BERT 剪成高稀疏模型然后在多个下游任务上继续 fine-tune并且保持稀疏模式不变同时精度损失很小这就是 “Prune Once for All” 的核心含义不是每个任务都重新剪枝而是先生成一个稀疏预训练语言模型然后让它服务于多个下游任务。三、核心思想Prune OFA 的核心思想是在预训练任务上结合权重剪枝和知识蒸馏得到一个高稀疏的预训练语言模型迁移到下游任务时不再重新剪枝而是锁定已有稀疏模式只 fine-tune 未被剪掉的权重。它的流程可以拆成三步。第一步准备 teacher。先得到一个在预训练数据上优化过的 dense teacher model。由于作者没有原始 BERT / DistilBERT 的完全相同预训练数据处理流程所以他们用自己处理的 English Wikipedia 数据先做一个 teacher preparation 步骤让 teacher 更适配后续剪枝所用数据。第二步剪 student。student 从 teacher 初始化然后在预训练任务损失和知识蒸馏损失的共同监督下训练同时使用Gradual Magnitude Pruning和Learning Rate Rewinding逐步剪掉低幅值权重。最后得到一个稀疏预训练模型。第三步下游迁移时锁定稀疏模式。稀疏预训练模型迁移到 SQuAD、MNLI、QQP、QNLI、SST-2 等任务时已经被剪成 0 的权重不会再恢复。这一步叫Pattern-lock也就是固定剪枝 mask只更新保留下来的非零权重。一句话说Prune OFA 不是在每个下游任务上重新搜索稀疏结构而是先训练一个通用稀疏 BERT再把这个稀疏 BERT 迁移到多个任务。四、它剪的是什么结构这篇论文做的是非结构化权重剪枝。也就是说它剪的是单个 weight而不是完整 attention head、FFN neuron、hidden dimension 或 Transformer layer。论文明确说它关注的是unstructured weight pruning。具体剪枝作用在 Transformer encoder 中的所有 Linear layers 上包括 pooler layer如果模型有 pooler 的话。所以剪枝后BERT 的层数不变。hidden size 不变。attention head 数量不变。FFN intermediate size 不变。矩阵 shape 不变。变化的是矩阵中大量权重被置为 0。这和 CoFi、Block Pruning、LayerDrop、Poor Man’s BERT 都不一样。Prune OFA 的目标不是直接得到一个结构更小的 dense BERT而是得到一个高稀疏的预训练 BERT。五、Gradual Magnitude Pruning 的作用Prune OFA 使用的是Gradual Magnitude PruningGMP。它的基本思想很简单训练过程中逐步提高稀疏率每隔一定步数把低幅值权重剪掉直到达到目标稀疏率。这个过程不是一开始就直接剪到 85% 或 90%。如果一开始就强行置零大量权重模型很容易崩。GMP 的好处是让模型有时间适应稀疏化过程。论文在方法部分说明它采用 Zhu 和 Gupta 的 gradual magnitude pruning 思路训练中每隔若干步剪掉当前最低幅值权重直到达到目标稀疏率。这里的重要点是Prune OFA 的重要性标准仍然是权重幅值。它不像 Movement Pruning 那样学习 movement score也不像 CoFi 那样学习 L0 mask。它采用的是更传统的 magnitude pruning但通过预训练阶段剪枝 蒸馏 学习率回绕 pattern lock把效果做得更强。六、Learning Rate Rewinding 的作用论文在 GMP 中加入了Learning Rate RewindingLRR。普通 gradual pruning 中学习率通常按原计划持续下降。但剪枝会不断改变模型结构如果学习率已经很低模型可能很难从剪枝损伤中恢复。LRR 的思想是每次剪枝后把学习率调度器回绕到剪枝开始时附近的状态让模型有更强的恢复能力。论文说它把 Renda 等人提出的 learning rate rewinding 思想结合进 GMP每隔若干步进行剪枝时把 learning rate scheduler 回绕到 pruning start 时刻的状态剪枝结束后再继续原始 schedule。消融实验也显示LRR 对结果有明显帮助。论文在附录中说使用 LRR 可以改善下游 fine-tuning 结果尤其在 SQuAD 上提升明显并且 LRR 与蒸馏结合效果更好。所以 LRR 在这里的作用可以理解为剪枝造成损伤后用更合适的学习率让稀疏模型恢复。七、知识蒸馏的作用Prune OFA 使用了两类蒸馏。第一类是预训练剪枝阶段的蒸馏。在 student pruning 阶段student 同时学习预训练任务目标和 teacher 的输出分布。也就是说student 不只是学习 masked language modeling / next sentence prediction 这类预训练任务还要模仿 dense teacher 的预测行为。论文中明确写到student pruning 阶段使用预训练任务损失和知识蒸馏损失的线性组合。第二类是下游任务迁移阶段的蒸馏。在 fine-tune 到具体任务时作者还可以为每个任务训练一个 dense task teacher然后让 sparse model 在任务上模仿它。论文结果表明迁移阶段加入知识蒸馏可以显著提升稀疏模型精度。这里有一个细节需要注意Prune OFA 的主要卖点是“无需每个任务重新剪枝”不是“下游阶段完全不需要蒸馏”。它可以不用下游蒸馏直接迁移但最强结果通常使用 transfer-time KD。也就是说稀疏结构只剪一次。但下游任务上仍然可以用 KD 来提升精度。八、Pattern-lock为什么很关键Pattern-lock 是这篇论文里一个很实用的工程机制。稀疏预训练模型迁移到下游任务时如果不做限制被剪成 0 的权重可能在 fine-tuning 中重新变成非零。这样虽然可能提升精度但稀疏模式就变了模型不再保持预训练阶段得到的稀疏结构。所以作者提出 Pattern-lock在下游 fine-tuning 时把已经为 0 的权重固定住不允许它们恢复。具体做法是先为每个稀疏层建立一个 mask非零位置为 1零位置为 0训练时mask 为 0 的位置梯度被置为 0因此这些权重始终保持 0。论文附录 B 对这个机制做了说明。Pattern-lock 的意义是保证“Prune Once”真的成立。否则模型在下游任务上训练一段时间后稀疏模式发生变化就变成另一种 task-specific 稀疏模型了。九、实验设置论文主要验证三种架构BERT-BaseBERT-LargeDistilBERT稀疏率设置为BERT-Base85%、90%DistilBERT85%、90%BERT-Large90%预训练剪枝使用 English Wikipedia约2500M words并划分为 95% train、5% validation。下游任务包括SQuAD v1.1和 GLUE 中的MNLI、QQP、QNLI、SST-2。论文给出了这些任务的数据量例如 SQuAD v1.1 约 89KMNLI 393KQQP 364KQNLI 105KSST-2 67K。这组实验的重点不是“某一个任务上最高分”而是验证同一个稀疏预训练模型能否迁移到多个不同任务并保持较小精度损失。十、主要结果解读10.1 BERT-Base 85% 稀疏几乎接近 dense BERTBERT-Base dense reference 在 SQuAD 上是80.80 EM / 88.50 F1。Prune OFA 85% 稀疏、不加下游 KD 时达到78.59 / 86.63加入下游 KD 后达到81.10 / 88.42几乎追平 dense BERT。MNLI、SST-2、QNLI、QQP 上也保持较小损失。这说明85% 稀疏的 BERT-Base在加入迁移阶段蒸馏后可以保持非常接近 dense BERT 的下游表现。更关键的是这个稀疏结构不是针对每个任务重新剪出来的而是预训练阶段统一得到的。10.2 BERT-Base 90% 稀疏仍然有竞争力在 90% 稀疏率下BERT-Base 的结果仍然不错。Prune OFA KD 在 SQuAD 上达到79.83 EM / 87.25 F1MNLI 是81.45 / 82.43SST-2 是90.88QNLI 是89.07QQP 是90.93 / 87.72。这说明即使只保留约 10% encoder linear 权重Prune OFA 仍然能保留相当多的迁移能力。不过和 85% 稀疏相比90% 稀疏下精度损失更明显。这个现象也符合前面 “Compressing BERT” 的结论低到中等剪枝率比较安全高稀疏下必须借助更强训练策略和蒸馏。10.3 BERT-Large 90% 稀疏大模型剪完后仍然很强BERT-Large dense reference 在 SQuAD 上是83.99 EM / 90.93 F1。Prune OFA 90% 稀疏后达到83.35 / 90.20QAT 后为83.22 / 90.02在 SST-2、QNLI、QQP 上也非常接近 dense reference。论文指出90% 稀疏 BERT-Large 的非零参数约30.2M并且其准确率优于 dense BERT-Base。这点非常重要大模型虽然原始参数多但冗余也多。剪到 90% 后BERT-Large 仍然可以保留很强的表示能力。这和 “train large, then compress” 的思想一致从大模型开始剪最后可能得到一个比小 dense 模型更强的稀疏模型。10.4 DistilBERT 也可以被 Prune OFA 继续压缩论文还测试了 DistilBERT。DistilBERT 本身已经是蒸馏压缩后的模型但 Prune OFA 仍然能把它剪到 85% 和 90% 稀疏。例如 85% 稀疏下DistilBERT Prune OFA 在 SQuAD 上达到78.10 / 85.82接近 dense DistilBERT 的77.70 / 85.8090% 稀疏下也有76.91 / 84.82。在多个任务上Prune OFA 比 fine-tune pruning baseline 更好。这说明即使是已经蒸馏过的小模型仍然存在大量可剪权重。不过 DistilBERT 的可压缩空间比 BERT-Large 更紧极高稀疏下精度更容易下降。10.5 和量化结合进一步压缩Prune OFA 还和Quantization-Aware TrainingQAT结合。论文摘要中给出的代表性结果是90% 稀疏 BERT-Large 在 SQuAD v1.1 上 fine-tune 后再做 8-bit QAT可以达到 encoder40× compression ratio同时精度损失小于 1%。论文结果部分还指出QAT 会带来额外平均约0.67% relative accuracy的精度下降但模型体积显著更小85% sparse QAT 的模型比 90% sparse full precision 模型更小体积约为后者的 0.375。这说明 Prune OFA 的实际压缩路线是先稀疏化再量化。也就是剪枝减少非零权重数量。量化减少每个非零权重的存储位宽。十一、和 Movement Pruning 的区别Movement Pruning 是task-specific pruning。它在下游任务 fine-tuning 过程中学习每个权重是否应该保留优势是对当前任务自适应强尤其高稀疏下比 magnitude pruning 更好。Prune OFA 的思路不同它不希望每个任务都重新剪。它希望先得到一个稀疏预训练模型然后多个任务共用这个稀疏结构。所以二者的区别可以概括为Movement Pruning每个任务剪一次任务适配更强。Prune OFA预训练阶段剪一次迁移多个任务更方便。论文也明确指出Movement Pruning 和类似方法通常需要对每个任务进行较长 fine-tuning并调整剪枝相关超参数Prune OFA 则避免了每个任务单独 pruning 和 tuning 的负担。十二、和 Compressing BERT 的关系前面你问过的Compressing BERT: Studying the Effects of Weight Pruning on Transfer Learning其实是 Prune OFA 的重要前置工作。Compressing BERT 的一个核心结论是BERT 在预训练阶段剪枝和在下游 fine-tuning 阶段剪枝对最终任务精度影响差别不大。Prune OFA 正是沿着这个结论往前推进既然预训练阶段剪枝可以迁移那就把预训练阶段剪枝做得更强剪到 85%–90% 稀疏并加入蒸馏、LRR 和 pattern-lock。论文也在引言和相关工作里明确引用 Gordon 等人的工作并说明自己是在其基础上提升到更高稀疏率和更好结果。所以这两篇论文的关系可以这样理解Compressing BERT 证明“预训练阶段剪枝可行”。Prune OFA 进一步证明“高稀疏预训练剪枝也可以迁移多个任务”。十三、和 CoFi / Block Pruning 的区别CoFi 和 Block Pruning 都更关注结构化加速。CoFi 会剪 MHA layer、FFN layer、attention head、FFN dimension、hidden dimension所以剪完后模型结构真的变小。Block Pruning 会通过 block、dimension、head 等较规则结构让稀疏更适合实际推理。Prune OFA 则不同它主要做非结构化权重稀疏。它的优势是高压缩率和迁移通用性。它的弱点是真实推理加速依赖稀疏硬件 / 稀疏 kernel 支持。论文虽然说稀疏神经网络可以减少计算和内存但因为它聚焦 unstructured pruning实际端到端 latency 不一定像结构化剪枝那样直接下降。所以如果从部署角度看Prune OFA 更偏存储压缩和稀疏模型路线。CoFi / Block Pruning 更偏结构化加速路线。十四、方法优点第一真正解决了“每个任务都要重新剪枝”的麻烦。Prune OFA 得到的是稀疏预训练模型不是某个任务专属稀疏模型。下游任务只需要 fine-tune 并保持 pattern-lock。第二高稀疏率下仍然能保持较好迁移性能。BERT-Base 可以做到 85% / 90% 稀疏BERT-Large 可以做到 90% 稀疏并且在多个任务上保持较小精度损失。第三结合了剪枝和蒸馏。预训练剪枝阶段使用 teacher distillation下游阶段还可以继续用 task teacher 做 distillation这让高稀疏模型更稳。第四可以和 8-bit QAT 叠加。剪枝减少非零参数量化降低单个参数位宽两者组合能得到很高 compression ratio。第五适用于多个架构。论文在 BERT-Base、BERT-Large、DistilBERT 上都做了实验说明方法不是只针对单一模型。十五、方法局限第一它是非结构化剪枝实际加速不一定直接。虽然参数和存储大幅下降但矩阵 shape 没变。如果没有高效稀疏矩阵乘法支持普通 GPU 上不一定能得到和稀疏率成比例的推理加速。第二预训练剪枝成本不低。它不是一个轻量后处理方法。论文中 Prune OFA 使用 English Wikipedia 做预训练任务并进行 100k steps 级别的训练设置这比下游任务剪枝更通用但前期成本更高。第三最佳结果仍然依赖下游 KD。Prune OFA 的稀疏结构可以迁移但表格中最强结果通常还使用 transfer-time knowledge distillation。也就是说它降低了“重新剪枝”的成本但并没有完全消除下游训练技巧。第四稀疏模式不是任务自适应最优。因为所有任务共享同一个预训练稀疏 mask所以它未必是某个具体任务的最优稀疏结构。Movement Pruning、CoFi 这类 task-specific 方法在某些任务上可能更有针对性。第五主要验证在 BERT 系列模型上。论文覆盖 BERT-Base、BERT-Large、DistilBERT但不能直接说明现代大规模 decoder-only LLM 上同样成立。LLM 的规模、训练目标、推理瓶颈和稀疏 kernel 支持都不同。十六、整体评价Prune Once for All 的核心贡献是把 BERT 剪枝从“任务级剪枝”推进到“预训练级稀疏模型”。它解决的问题很实际如果每个任务都要重新剪枝那么剪枝的工程成本很高也不利于发布通用压缩模型。Prune OFA 证明可以先得到一个高稀疏预训练语言模型再把它迁移到多个下游任务并通过 pattern-lock 保持稀疏结构。它的逻辑链条很清楚先在预训练任务上用 GMP LRR KD 得到稀疏预训练模型。再在下游任务中保持稀疏 pattern 不变。必要时加入 task-level KD。最后还可以叠加 8-bit QAT。从剪枝谱系看它处在一个很重要的位置Compressing BERT 证明预训练剪枝可行。Prune OFA 把预训练剪枝做到高稀疏和多任务迁移。Movement Pruning 更强调任务自适应高稀疏。Block Pruning 和 CoFi 则进一步追求结构化加速。所以这篇论文最值得记住的不是某一个复杂公式而是这个思想稀疏模式可以在预训练阶段一次性学习然后迁移到多个下游任务。十七、一句话总结《Prune Once for All: Sparse Pre-Trained Language Models》提出 Prune OFA通过在预训练阶段结合 Gradual Magnitude Pruning、Learning Rate Rewinding 和知识蒸馏训练出 85%–90% 稀疏的 BERT-Base、BERT-Large 和 DistilBERT这些稀疏预训练模型在下游任务 fine-tuning 时通过 Pattern-lock 保持稀疏结构不变从而避免每个任务重新剪枝。它的核心价值是证明“高稀疏预训练语言模型”可以作为通用压缩模型迁移到多个 NLP 任务并且还能与下游蒸馏和 8-bit 量化进一步组合。