Sheared LLaMA: Accelerating Language Model Pre-training via Structured Pruning 解读 一、论文基本信息论文题目Sheared LLaMA: Accelerating Language Model Pre-training via Structured Pruning方法名称LLM-Shearing作者Mengzhou Xia、Tianyu Gao、Zhiyuan Zeng、Danqi Chen发表ICLR 2024官方代码仓库为princeton-nlp/LLM-Shearing仓库提供了 Sheared-LLaMA 的剪枝与 continued pre-training 代码以及 1.3B、2.7B 模型权重。(GitHub)一句话先概括Sheared LLaMA 不是像 SparseGPT / Wanda 那样做非结构化权重稀疏也不是像 LLM-Pruner 那样剪完后用 LoRA 恢复而是从一个强大的大模型 LLaMA2-7B 出发通过“目标结构化剪枝 继续预训练”低成本生产出强性能的小规模 base LLM。它的核心目标不是简单让已有模型变稀疏而是回答一个更大的问题能不能不从零训练 1B / 3B 小模型而是直接从已有强大 7B 模型中“剪出”一个小模型再用少量 token 继续预训练使其超过同规模从头训练模型二、这篇论文要解决什么问题训练小规模 LLM 也很贵。比如 1B、3B 这种模型虽然比 7B、13B 小但如果从零预训练仍然需要几百 B 到 1T 级别 token。论文指出训练每一个不同规模的开放 LLM 都要消耗大量计算资源因此作者提出的问题是能否利用已有大模型以更少计算得到一个通用、强性能的小模型传统 LLM 剪枝大多有两个方向第一剪完直接用。例如 SparseGPT、Wanda主要做非结构化权重剪枝剪完后尽量不训练。第二剪完用少量任务数据恢复。例如 LLM-Pruner、LoRAPrune更多是结构化剪枝 LoRA 或轻量恢复。Sheared LLaMA 的定位不同。它不是只想在已有大模型上“省一点推理成本”而是想把剪枝变成一种生产小型 base model 的预训练加速方法。所以它的核心问题是与其从零训练一个 1.3B / 2.7B 模型能不能先把 LLaMA2-7B 结构化剪成目标大小再继续预训练少量 token得到更强的小模型三、核心思想Sheared LLaMA 的核心方法叫LLM-Shearing包含两个关键技术第一Targeted Structured Pruning。把大模型剪到一个预先指定的目标结构例如指定层数、hidden dimension、attention head 数、FFN intermediate dimension。论文明确说该方法会通过删除layers、heads、intermediate dimensions、hidden dimensions把大模型端到端剪到目标形状。(arXiv)第二Dynamic Batch Loading。剪枝后继续预训练时不再固定使用原始数据比例而是根据不同数据域的 loss 恢复速度动态调整 batch 中各个 domain 的采样比例。论文指出剪枝模型在不同数据域中保留知识的程度不同因此继续预训练时应该给恢复慢的 domain 更多数据。也就是说Sheared LLaMA 不是单纯“剪模型”而是先剪出一个目标形状的小模型。再用更聪明的数据采样方式继续预训练。最终得到一个真正可用的小型 base LLM。四、它剪的是什么Sheared LLaMA 是明确的结构化剪枝。它剪的包括Transformer layers。Hidden dimensions。Attention heads。FFN intermediate dimensions。论文方法部分说明它为不同粒度引入 pruning masks包括全局的 layers 和 hidden dimensions以及局部的 attention heads 和 intermediate dimensions每个 mask 控制对应子结构是保留还是删除。(ar5iv)它不是非结构化权重剪枝。N:M 半结构化稀疏。token pruning。KV cache pruning。单纯 layer dropping。所以如果放到 LLM 剪枝分类里它属于targeted structured pruning continued pre-training。五、为什么叫 Targeted Structured Pruning很多结构化剪枝方法的问题是剪完之后结构可能很不规则。例如不同层 head 数不一样。不同层 FFN 宽度不一样。hidden dimension 可能不符合常见硬件友好配置。这种不规则结构理论上参数少但推理时可能不好部署甚至带来额外 overhead。论文就指出已有结构化剪枝方法可能产生偏离常见架构的不规则配置从而影响推理效率。(ar5iv)Sheared LLaMA 的做法是不是只给一个稀疏率而是给一个目标模型形状。例如我要把 LLaMA2-7B 剪成类似 1.3B 模型的结构。我要把 LLaMA2-7B 剪成类似 2.7B 模型的结构。论文中作者用Pythia-1.4B的结构作为 1.3B 目标结构用INCITE-Base-3B的结构作为 2.7B 目标结构。这就是 “targeted” 的含义目标不是任意剪小而是剪成一个预先设定、推理友好、接近标准小模型配置的 dense architecture。六、剪枝 mask 是怎么学的Sheared LLaMA 借鉴了 CoFiPruning / L0 regularization 这类方法。它给不同结构单元加上可学习 mask并用 hard concrete distribution 让 mask 接近 0 或 1。论文明确说这些 mask 通过 hard concrete distributions 参数化可以集中到 0 或 1从而对应离散的剪枝 / 保留决策。(ar5iv)简单理解mask 接近 1保留这个 layer/head/channel/neuron。mask 接近 0删除这个结构。训练时同时优化语言模型 loss。结构约束。这里的结构约束不是“总体剪掉多少参数”而是“最终结构要符合目标模型形状”。论文用 Lagrange multipliers 来约束目标层数、目标 hidden dimension、目标 head 数和目标 intermediate dimension。所以它不是简单按重要性排序一次性删除而是在一个 constrained optimization 里学习哪些层保留。哪些 head 保留。哪些 hidden dimensions 保留。哪些 FFN intermediate dimensions 保留。最后再把 mask 接近 0 的结构物理删除得到目标小模型。七、为什么还要 continued pre-training这是 Sheared LLaMA 和 LLM-Pruner / LoRAPrune 的重要区别。很多剪枝论文的流程是剪枝 → 少量恢复训练 / LoRA 微调 → 评估。Sheared LLaMA 认为对于通用 base LLM这远远不够。结构化剪枝一定会损失语言建模能力如果想得到一个真正强的小模型必须进行continued pre-training。论文明确把流程分成两阶段第一阶段把 source model 剪成 target model。第二阶段继续用语言建模目标预训练剪枝模型。作者还强调后者对生产竞争力小模型至关重要。所以它的目标不是“剪完尽量不掉”而是用剪枝获得一个强初始化再用少量继续预训练把能力恢复甚至提升。这也解释了为什么论文标题说Accelerating Language Model Pre-training它把剪枝当成加速预训练的一种方式而不是单纯推理压缩技巧。八、Dynamic Batch Loading为什么重要LLM 预训练数据通常来自多个 domain例如CommonCrawl。C4。GitHub。Wikipedia。Books。ArXiv。StackExchange。Sheared LLaMA 使用 RedPajama 数据并将其划分为这些 domain论文说明每个 domain 都构建了 held-out validation set。作者观察到剪枝模型在不同 domain 上恢复速度不一样。例如有些低熵、小规模 domain 中的知识可能在剪枝后保留得更多而 C4 这种高熵、大规模 domain 上的能力恢复可能更慢。论文给出的解释是剪枝模型在不同 domain 中保留的知识量不同如果继续按原始预训练比例采样就会浪费数据恢复效率低。因此 Dynamic Batch Loading 的思路是定期评估每个 domain 的 validation loss。看它距离目标 reference loss 还有多远。哪个 domain 恢复慢就在后续 batch 中采样更多。哪个 domain 已经恢复得好就减少采样。这比固定数据配比更适合剪枝后的模型因为剪枝后的模型不是一个随机初始化小模型它已经从大模型继承了一部分 domain knowledge。九、它和“从头训练小模型”有什么区别从头训练 2.7B 模型的流程是随机初始化 2.7B 模型。用几百 B 到 1T token 预训练。慢慢学语言、知识、推理模式。Sheared LLaMA 的流程是从 LLaMA2-7B 继承权重和知识。结构化剪成 2.7B。只用 50B tokens 继续预训练。论文和项目页都强调Sheared-LLaMA-1.3B 和 2.7B 是从 LLaMA2-7B 剪出来后只训练了 50B tokens项目页还称这相当于之前强开源 3B 模型训练预算的 5%。这就是它的主要观点强大大模型本身就是一个很好的初始化。剪枝后的模型虽然一开始掉性能但它恢复得很快。用少量 continued pre-training 就能超过很多从头训练的小模型。十、主要实验设置源模型是LLaMA2-7B。作者将其剪成两个目标规模Sheared-LLaMA-1.3B。Sheared-LLaMA-2.7B。训练数据使用RedPajama因为 LLaMA2 的原始训练数据不可公开论文中剪枝阶段使用约0.4B tokenscontinued pre-training 使用50B tokens序列长度保持 LLaMA2 风格的4096。(ar5iv)评估任务包括 commonsense、reading comprehension、world knowledge、MMLU、NQ 等。论文使用 lm-evaluation-harness并报告 zero-shot / few-shot 指标。十一、主要结果Sheared-LLaMA-1.3B 的平均下游表现为51.0超过 OPT-1.3B 的48.2和 Pythia-1.4B 的48.9。Sheared-LLaMA-2.7B 的平均下游表现为56.7超过 OPT-2.7B、Pythia-2.8B、INCITE-Base-3B、OpenLLaMA-3B-v1 和 OpenLLaMA-3B-v2 等同规模模型。项目页也总结说Sheared-LLaMA-2.7B 在同规模开源模型中表现更好并且只用了约3%的 compute 达到与 OpenLLaMA-3B-v2 相当或更强的效果。这个结果非常关键。它说明结构化剪枝不是只能做任务特定压缩。如果剪枝后继续预训练得当它可以成为生产强小模型的路线。十二、Sheared LLaMA 和 LLM-Pruner 的区别两者都是结构化剪枝但目标完全不同。LLM-Pruner更像是把已有 LLM 压缩一下用少量数据和 LoRA 恢复目标是 task-agnostic compression。Sheared LLaMA更像是“从大模型生产小模型”的预训练路线。它不满足于剪完后少量恢复而是继续预训练 50B tokens目标是得到强 base LLM。可以这样理解LLM-Pruner 是压缩已有模型。Sheared LLaMA 是加速小模型预训练。LLM-Pruner 更适合资源有限、想快速得到一个稍小模型的场景Sheared LLaMA 更适合有一定预训练预算、想生产高质量 1B/3B base model 的场景。十三、和 LoRAPrune 的区别LoRAPrune 把 LoRA 和结构化剪枝结合用 LoRA 梯度指导剪枝并通过 LoRA 微调恢复。它关注的是PEFT-aware structured pruning。Sheared LLaMA 不依赖 LoRA 作为核心恢复机制而是继续用语言建模目标做 continued pre-training。它的目标更接近剪枝后的模型继续作为 base model 训练。所以区别是LoRAPrune剪枝 LoRA 微调强调低显存和推理结构变小。Sheared LLaMA剪枝 continued pre-training强调低成本产生强小型 base LLM。如果你要做下游任务压缩LoRAPrune 更直接如果你要生产一个通用小语言模型Sheared LLaMA 的路线更合适。十四、和 SparseGPT / Wanda 的区别SparseGPT 和 Wanda 默认是非结构化剪枝。它们主要是把权重矩阵里某些元素置零。不改变 hidden size、层数、head 数。剪完尽量不用训练。Sheared LLaMA 是结构化剪枝。它会删除层。head。hidden dimensions。FFN intermediate dimensions。并且最终得到的是一个标准小 dense 模型而不是稀疏大矩阵。论文也明确说 targeted structured pruning 会把模型剪成 specified target shape。所以它和 SparseGPT / Wanda 的根本区别是SparseGPT / Wanda 主要是 post-training sparsification。Sheared LLaMA 是 structural model resizing continued pre-training。十五、为什么 Sheared LLaMA 更像“模型生产方法”这篇论文最值得注意的是它不把剪枝看作模型压缩的终点而是把剪枝看作一个新的起点。传统剪枝通常问剪完还能保留多少原模型性能Sheared LLaMA 问的是剪出来的小模型继续训练后能不能比同规模从头训练模型更强这两个问题很不同。它证明了一个很有价值的方向已有强大模型可以作为“小模型预训练”的母模型。先剪成目标结构再继续预训练可能比从头训练更高效。这对开源小模型生产很重要。因为很多机构可能没有预算从零训练 1T tokens但如果能从已有 7B / 13B / 更大模型剪出目标规模再用几十 B tokens 继续训练就能更低成本得到强模型。十六、它是不是剪枝是的Sheared LLaMA 是结构化剪枝。但它不是普通意义上的“剪完就部署”的剪枝论文。更准确地说它是targeted structured pruning for LLM pre-training acceleration。它既是剪枝方法也是小模型生产路线。分类上可以写成LLM structured pruning。Model shearing / model resizing。Prune-and-continue-pretrain。Pretraining-efficient small LLM construction。十七、方法优点第一剪完后是标准 dense 小模型。它不是非结构化稀疏矩阵因此更容易用普通推理框架部署。第二目标结构推理友好。它不是随意剪成不规则结构而是对齐预先指定的目标架构避免不规则剪枝造成推理 overhead。第三性能强。Sheared-LLaMA-2.7B 只用 50B tokens 继续训练就超过多个同规模、训练 token 更多的开源模型。第四动态数据配比很有价值。Dynamic Batch Loading 针对剪枝模型不同 domain 恢复速度不一致的问题提升 continued pre-training 的数据效率。第五证明了剪枝可以服务于预训练。它把剪枝从“压缩模型”推进到“高效生产小模型”。十八、方法局限第一不是 training-free。它需要继续预训练 50B tokens。虽然比从头训练少很多但仍然不是普通实验室随便能跑的小成本。第二依赖强 source model。如果源模型不够强剪出来的小模型也未必有优势。项目页也总结初始 base model 越强得到的 pruned model 越强。第三主要实验从 LLaMA2-7B 剪到 1.3B / 2.7B。论文说方法可以扩展到更大模型但主实验仍集中在 7B 源模型。第四剪枝阶段本身较慢。论文提到 pruning stage 比标准 LM training 慢很多因此实际只给剪枝阶段较有限预算然后再继续预训练。第五需要预训练数据和训练系统支持。Dynamic Batch Loading、长序列 4096、RedPajama 多域数据、Composer / FlashAttention 等工程栈对复现有一定要求。十九、整体评价Sheared LLaMA 是 LLM 结构化剪枝方向中非常有代表性的一篇论文因为它把剪枝的目标从“压缩已有模型”转向“低成本生产强小模型”。它的核心观点可以概括为不要从零训练小模型。先从强大模型中剪出目标结构。再用少量 continued pre-training 恢复能力。训练数据还要根据剪枝模型的 domain 恢复情况动态调整。如果把它放到你最近看的 LLM 剪枝脉络中SparseGPT非结构化二阶重构one-shot。Wanda非结构化权重 × 激活one-shot。LLM-Pruner结构化dependency TaylorLoRA 恢复。LoRAPrune结构化LoRA-guided criterionPEFT-aware recovery。Sheared LLaMA结构化target shape pruning continued pre-training用剪枝加速小型 base model 生产。所以它最准确的位置是pretraining-oriented structured pruning for LLMs。二十、一句话总结《Sheared LLaMA: Accelerating Language Model Pre-training via Structured Pruning》提出 LLM-Shearing先用 targeted structured pruning 将 LLaMA2-7B 剪成预先指定的 1.3B / 2.7B 目标结构剪枝对象包括 layers、hidden dimensions、attention heads 和 FFN intermediate dimensions再用 dynamic batch loading 进行 continued pre-training根据不同数据域 loss 恢复速度动态调整采样比例。它不是 SparseGPT / Wanda 式非结构化 one-shot 剪枝而是“结构化剪枝 继续预训练”的小模型生产路线实验表明Sheared-LLaMA-1.3B 和 2.7B 只用 50B tokens 就能超过多个同规模开源模型说明从强大 LLM 中剪出小模型再继续训练是比从零预训练更高效的一条路线。