Back to blog

EfficientLLM: Scalable Pruning-Aware Pretraining for Architecture-Agnostic Edge Language Models

通过剪枝感知预训练实现架构无关的高效边缘语言模型,首次将LLM压缩提升到预训练阶段

EfficientLLM: Scalable Pruning-Aware Pretraining for Architecture-Agnostic Edge Language Models

一、论文概述

项目内容
标题EfficientLLM: Scalable Pruning-Aware Pretraining for Architecture-Agnostic Edge Language Models
作者Xingrun Xing, Zheng Liu, Shitao Xiao, Boyan Gao, Yiming Liang, Wanpeng Zhang, Haokun Lin, Guoqi Li, Jiajun Zhang
机构中国科学院自动化研究所 (CASIA)、北京人工智能研究院 (BAAI)
论文https://arxiv.org/abs/2502.06663
代码https://github.com/Xingrun-Xing2/EfficientLLM
HuggingFacehttps://huggingface.co/collections/xrxing/efficientllm-pruning-aware-pretraining-67a8ecc6a49580b647a6184f
发布2025-02-10 (v1), 2025-02-11 (v2)
许可MIT License
领域Machine Learning (cs.LG), ICML

二、核心思想

问题定义

现代大语言模型(LLMs)依靠缩放定律(scaling laws)在大规模参数下实现智能涌现。然而,对于边缘设备部署,存在以下关键挑战:

  1. 云端成本高:典型LLM(7B-1T参数)需要云端部署和持续网络连接
  2. 延迟问题:移动边缘应用对实时性要求高
  3. 隐私顾虑:数据需要上传到云端处理

现有方法的局限:

  • 直接预训练:受限于缩放定律,小模型的数据效率低下,且智能涌现仅在较大模型规模出现
  • 后训练剪枝:仅使用小型校准数据集,导致显著性能退化
  • 知识蒸馏:教师模型(通常7B)的计算量是边缘模型的50倍以上

解决方案概述

本文提出剪枝感知预训练(Pruning-Aware Pretraining),核心思想是在预训练阶段持续进行结构化剪枝,而非传统的后训练剪枝。这具有两个关键特性:

  1. 数据可扩展(Data-scalable):在LLM中引入最小参数组,持续优化结构化剪枝,将LLM-Pruner和SparseGPT等后训练剪枝方法扩展到预训练阶段
  2. 架构无关(Architecture-agnostic):使用显著性驱动的剪枝自动设计LLM架构,首次在现代预训练中超越人类设计的最优架构

与现有方法的对比:

  • 相比直接预训练:利用更大优化模型的性能,小模型永远无法通过单独预训练达到
  • 相比后训练剪枝:用预训练数据扩展剪枝阶段,显著提升压缩质量

三、技术架构

整体框架

剪枝感知预训练被形式化为一个双层优化问题:

外层优化:选择最优的剪枝mini-group g*
内层优化:更新模型权重 w*

训练循环:
1. 梯度下降步骤(权重更新)
2. mini-group优化步骤(剪枝决策)
3. 二阶权重更新(补偿剪枝误差)

核心公式

问题形式化

给定一个优化的大模型 M,剪枝后的模型 M* 可以表示为:

M∗=M−∑t=1ngt,s.t.min⁡gt∈GLpretrain(M)\mathcal{M}^* = \mathcal{M} - \sum_{t=1}^{n} g_t, \quad \text{s.t.} \quad \min_{g_t \in \mathcal{G}} \mathcal{L}_{pretrain}(\mathcal{M})

其中 gtg_t 是第t步剪枝的mini-group参数,G\mathcal{G} 是由mini-groups构成的剪枝空间。

递归剪枝

剪枝过程被解耦为t步,可以近似顺序求解:

Mt=Mt−1−gt∗\mathcal{M}_t = \mathcal{M}_{t-1} - g_t^*

gt∗=argmingt∈GLpretrain(gt∣Mt−1)g_t^* = \underset{g_t \in \mathcal{G}}{\mathrm{argmin}} \mathcal{L}_{pretrain}(g_t | \mathcal{M}_{t-1})

双层优化

将问题转化为mini-groups g和权重w的双层优化:

min⁡g∈GLpretrain(g,w∗∣M)\min_{g \in \mathcal{G}} \mathcal{L}_{pretrain}(g, w^* | \mathcal{M}) s.t.w∗=argminwLpretrain(w,g∗∣M)\text{s.t.} \quad w^* = \underset{w}{\mathrm{argmin}} \mathcal{L}_{pretrain}(w, g^* | \mathcal{M})

外层优化通过Eq.3求解,内层优化通过梯度下降直接求解。梯度下降和mini-group优化交替进行,称为剪枝感知预训练 x1。

二阶Taylor展开

对于优化后的模型,任何权重 w 的损失可以用二阶Taylor展开近似:

L(w)≃L(w∗)+δw⊤∇L(w∗)+12δw⊤HL(w∗)δw\mathcal{L}(\mathbf{w}) \simeq \mathcal{L}(\mathbf{w}^*) + \delta\mathbf{w}^\top \nabla\mathcal{L}(\mathbf{w}^*) + \frac{1}{2}\delta\mathbf{w}^\top \mathbf{H}_{\mathcal{L}}(\mathbf{w}^*) \delta\mathbf{w}

其中 δw=w−w∗\delta\mathbf{w} = \mathbf{w} - \mathbf{w}^*,HL\mathbf{H}_{\mathcal{L}} 是Hessian矩阵。

显著性计算

最优mini-group的选择通过显著性评估:

gt∗=argmingt∈G{Sattn,Sffn,Sstem}g_t^* = \underset{g_t \in \mathcal{G}}{\mathrm{argmin}} \{ \mathcal{S}_{attn}, \mathcal{S}_{ffn}, \mathcal{S}_{stem} \}

其中三种剪枝类型的显著性分别为:

  • 类型I (注意力头剪枝):Sattn\mathcal{S}_{attn} - 按行求和的注意力输出投影显著性
  • 类型II (FFN通道剪枝):Sffn\mathcal{S}_{ffn} - 按行求和的down投影显著性
  • 类型III (Stem通道剪枝):Sstem\mathcal{S}_{stem} - 按列求和的输出层组显著性

二阶权重更新

剪枝后剩余权重的更新公式:

δwp=−wp[H−1]pp⋅H:,p−1\delta w_p = -\frac{w_p}{[\mathbf{H}^{-1}]_{pp}} \cdot \mathbf{H}^{-1}_{:,p}

为高效计算Hessian逆,通过求解线性方程: ep=HH:,p−1\mathbf{e}_p = \mathbf{H} \mathbf{H}^{-1}_{:,p}

模型组件

最小剪枝组(Minimal Pruning Groups)

定义三种基本剪枝类型:

类型名称剪枝单元耦合参数
类型I注意力头剪枝每个attention headQ/K/V输入通道 + O输出通道
类型IIFFN通道剪枝每个FFN中间通道Up/Gate输入通道 + Down输出通道
类型IIIStem通道剪枝每个transformer stem通道Embedding + 所有层Q/K/V/O + FFN + LM Head

公式表示:

类型I - 注意力mini-group: Gattn={W:,i:j(k,ℓ),W:,i:j(q,ℓ),W:,i:j(v,ℓ),Wi:j,:(o,ℓ),ℓ=1,2,...,n}\mathcal{G}_{attn} = \{W_{:,i:j}^{(k,\ell)}, W_{:,i:j}^{(q,\ell)}, W_{:,i:j}^{(v,\ell)}, W_{i:j,:}^{(o,\ell)}, \ell=1,2,...,n\}

类型II - FFN mini-group: Gffn={W:,i(up,ℓ),W:,i(gate,ℓ),Wi,:(down,ℓ),ℓ=1,2,...,n}\mathcal{G}_{ffn} = \{W_{:,i}^{(up,\ell)}, W_{:,i}^{(gate,\ell)}, W_{i,:}^{(down,\ell)}, \ell=1,2,...,n\}

类型III - Stem mini-group: Gstem={Wi,:(k,ℓ),Wi,:(q,ℓ),Wi,:(v,ℓ),W:,i(o,ℓ)}∪{Wi,:(up,ℓ),Wi,:(gate,ℓ),W:,i(down,ℓ)}∪{wi(emb),Wi,:(head)}\mathcal{G}_{stem} = \{W_{i,:}^{(k,\ell)}, W_{i,:}^{(q,\ell)}, W_{i,:}^{(v,\ell)}, W_{:,i}^{(o,\ell)}\} \cup \{W_{i,:}^{(up,\ell)}, W_{i,:}^{(gate,\ell)}, W_{:,i}^{(down,\ell)}\} \cup \{\mathbf{w}_i^{(emb)}, W_{i,:}^{(head)}\}

剪枝空间

原始剪枝空间:h(ℓ)×n(ℓ)×mh^{(\ell)} \times n^{(\ell)} \times m

通过mini-group耦合后,每步选择空间降至3,总剪枝空间为 3t3^t。

EfficientLLM变体

变体说明特点
EfficientLLM-A基础版本使用LLM-Pruner近似Eq.9
EfficientLLM-B增强版本在A基础上增加二阶权重更新

训练流程

训练阶段

  1. 剪枝感知预训练阶段:

    • 持续进行结构化剪枝和权重优化
    • 每步:梯度下降 + mini-group选择 + 二阶权重更新
    • 目标:从源模型剪枝到目标模型大小
  2. 继续预训练阶段:

    • 达到目标大小后,继续预训练以提升性能
    • 50B-500B tokens

训练超参数

参数值
批量大小1M tokens
GPU数量32-64 A800 GPUs
模型规模100M - 1.1B parameters
剪枝比例高达70%+

各模型训练详情

模型源模型剪枝预训练tokens继续预训练tokens
EfficientLLM-134MSmolLM-360M50.3B500B
EfficientLLM-469MSmolLM-1.7B72.1B500B
EfficientLLM-1.1BSmolLM-1.7B36.7B320B

数据组成

与源模型SmolLM保持相似的数据分布:

  • FineWeb-Edu: 220B tokens
  • Cosmopedia v2: 28B tokens
  • Python-Edu: 4B tokens
  • OpenWebMath: 27.5B tokens (随机采样)

Hessian近似策略

  1. 显著性检测:使用全局对角Hessian矩阵近似(如LLM-Pruner)
  2. 权重更新:使用逐层近似 HL≃XXT\mathbf{H}_{\mathcal{L}} \simeq XX^T

这种解耦策略兼顾了全局显著性检测和局部误差最小化。

四、核心创新

创新点说明理论/实验依据
剪枝感知预训练首次将LLM压缩提升到预训练阶段,用预训练数据扩展剪枝即使使用vanilla LLM-Pruner指标,也能超越SoTA剪枝方法
架构无关设计使用显著性驱动的剪枝自动设计LLM架构,首次超越人类设计自动发现的架构与人类最佳实践(如MobileLLM deep-and-thin)竞争
最小参数组定义三种最小剪枝单元,将剪枝空间从指数级降至线性原始空间 h×n×mh \times n \times m 降至每步选择3种类型
高效二阶更新解耦Hessian矩阵在显著性检测和权重更新中的应用显著性用全局对角Hessian,权重更新用逐层近似
数据高效预训练50B tokens超越17T tokens训练的Qwen2.5-0.5BEfficientLLM-469M仅用50B tokens超越SmolLM-360M(600B tokens)
高比例剪枝实现70%+剪枝比例且保持合理性能在Llama-7B上70%剪枝仍保持51.33%平均准确率

五、代码实现分析

仓库结构

GitHub仓库目前为初始版本,包含:

  • README.md - 项目说明和使用示例
  • LICENSE - MIT许可证
  • imgs/ - 论文图片

使用方式

通过HuggingFace Transformers加载预训练模型:

from transformers import AutoModelForCausalLM, AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("xrxing/EfficientLLM-469M", use_fast=False)

model = AutoModelForCausalLM.from_pretrained(
    "xrxing/EfficientLLM-469M",
    trust_remote_code=True,
    attn_implementation="flash_attention_2"
)

HuggingFace模型

模型参数量源模型继续预训练tokens
EfficientLLM-134M134MSmolLM-360M500B
EfficientLLM-469M469MSmolLM-1.7B50B / 500B
EfficientLLM-1.1B1.1BSmolLM-1.7B50B / 320B

待完成

  • 发布技术报告
  • 发布HuggingFace模型
  • 评估代码
  • 预训练代码
  • 演示和应用

六、实验结果

基准测试

评估任务

  • 世界知识:MMLU
  • 常识推理:ARC-c, ARC-e, BoolQ, HellaSwag, OBQA, PIQA, WinoGrande

主要结果 (Table 2)

100M-200M参数规模

模型#Tokens#ParamsMMLUARC-cARC-eBoolQHellaSwagOBQAPIQAWinoGrandeAvg.
OPT-125M180B125M26.0222.8743.3155.4431.3727.8062.6249.8041.89
GPT-neo-125M300B125M26.8923.2943.2261.7730.4926.0062.6251.9342.76
Pythia-160M300B162M26.4322.2737.8443.3329.9726.4058.8749.9638.38
Memba-130M1.2T130M27.6524.4947.5654.6835.1129.0064.6953.3544.13
MobileLLM-LS-125M1T125M-28.745.860.439.541.165.752.147.61
SmolLM-135M600B135M30.0529.3561.3259.8542.6734.4068.5552.9649.87
EfficientLLM-A500B134M30.5430.9762.8860.4043.8133.6068.8253.2850.54

300M-500M参数规模

模型#Tokens#ParamsMMLUARC-cARC-eBoolQHellaSwagOBQAPIQAWinoGrandeAvg.
OPT-350M180B331M26.9623.9844.0257.8036.6327.8064.9152.9644.01
BLOOM-560M350B559M27.3224.4046.0444.4636.5428.8062.5753.2042.29
Pythia-410M300B405M29.1024.1551.3959.2040.2029.4066.7053.8346.41
MobileLLM-LS-350M1T345M-32.554.462.850.645.869.857.253.30
SmolLM-360M600B362M33.8936.2670.1655.2353.5137.6071.3857.2254.48
Qwen2-0.5B15T494M31.8528.5055.0561.2549.1632.8069.7557.2250.53
Qwen2.5-0.5B17T494M33.3732.1764.4461.9952.0935.2070.2956.2053.20
EfficientLLM-A50B469M33.0935.9270.5059.8553.1635.0072.6956.2754.77
EfficientLLM-A500B469M34.5438.4072.1062.4256.8440.4073.8357.4657.35

1B参数规模

模型#Tokens#ParamsMMLUARC-cARC-eBoolQHellaSwagOBQAPIQAWinoGrandeAvg.
OPT-1.3B180B1.3B29.5730.0357.4956.5453.6632.8072.3159.0451.70
GPT-neo-1.3B380B1.3B30.0025.9456.3161.9048.9933.4071.0054.6250.31
BLOOM-1.1B350B1.1B29.1625.7751.7359.5143.1129.6067.3054.6247.38
Pythia-1B300B1.0B30.1426.9656.8660.0447.1531.2070.2952.8849.34
TinyLlama-1.1B3T1.1B32.3030.2960.4056.8559.1335.8073.0759.0453.51
ShearedLlama-1.3B50B1.3B31.5129.4461.0761.8359.3334.4073.9458.0154.00
OLMo-1B2T1.2B32.0330.7263.5561.3862.8636.4075.3559.3555.66
Llama3.2-1B-1.2B36.3131.4865.2863.8863.6937.4074.5960.5456.69
EfficientLLM-A50B1.1B36.7140.3673.6162.3960.2440.2075.1961.2559.03
EfficientLLM-A320B1.1B37.7142.2473.4867.0964.0941.8075.4161.1760.75

与LLM剪枝方法对比 (Llama-7B)

50%剪枝比例

模型#TuningARC-cARC-eBoolQHellaSwagOBQAPIQAWinoGrandeAvg.
MaP✓30.6349.3239.6942.4931.4066.8150.6744.43
MvP✓26.7944.0759.9440.9831.8063.0655.6446.04
WANDA✓34.2042.6850.9038.1238.7857.3855.9845.43
LLM-Pruner✓28.2446.4661.4747.5635.2068.8255.0948.98
LoRAPrune✓31.6245.1361.8847.8634.9871.5355.0149.72
LoRAShear✓32.2647.6862.1248.0134.6171.8056.2950.40
Compresso✓27.8248.8260.0939.3133.4066.7051.9346.87
NutePrune✗31.7446.5962.2053.8735.8069.9157.7751.13
NutePrune✓32.1751.6862.2655.8834.4071.0057.5452.13
EfficientLLM-A✗30.8052.1562.2954.7035.2071.3356.7551.89
EfficientLLM-A✓34.0464.8164.8360.1234.6073.8861.4856.25

70%剪枝比例

模型#TuningARC-cARC-eBoolQHellaSwagOBQAPIQAWinoGrandeAvg.
LLM-Pruner✓24.8339.5647.2831.6628.8060.8350.7540.53
NutePrune✓26.1942.1762.0839.4330.2062.3051.4644.83
EfficientLLM-A✗27.7354.5047.8947.7731.0068.1755.1747.46
EfficientLLM-A✓29.9558.5958.1352.0234.6070.0855.9651.33

与ShearedLlama对比 (Llama2-7B)

模型#Pruning#TuningARC-cARC-eBoolQHellaSwagOBQAPIQAWinoGrandeAvg.
ShearedLlama-2.7B0.4B–26.3749.6259.0247.1533.0066.5953.8347.94
EfficientLLM-A-2.7B0.66B–26.7155.5157.3147.5832.2067.1456.5148.99
EfficientLLM-A-2.7B10.56B–29.6161.1562.2354.9134.4071.1157.4652.98
ShearedLlama-1.3B0.4B–22.7841.0860.1834.6628.2063.0050.6742.94
EfficientLLM-B-1.3B0.95B–24.5745.7560.4038.2430.4063.1151.5444.86
ShearedLlama-1.3B0.4B5.6B26.5455.6060.1248.2831.6068.7756.2049.59
EfficientLLM-B-1.3B10.56B5B30.2060.0260.1853.2231.6070.4058.0151.95

消融实验

可扩展性分析

设置剪枝步骤与梯度下降步骤的比例为4:1, 2:1, 1:1, 1:9:

比例达到目标模型大小所需tokens
4:12.5B
2:14.5B
1:18.4B
1:972.1B

结论:扩展剪枝感知预training持续提升剪枝性能,通过在预训练阶段扩展LLM剪枝,可以扩展LLM压缩的上界。

泛化性分析

EfficientLLM框架可泛化到不同的二阶Taylor展开指标:

  • EfficientLLM-A:使用LLM-Pruner指标
  • EfficientLLM-B:使用SparseGPT指标 + 二阶权重更新

发现:

  • 在大规模剪枝感知预训练(>1B tokens)中,A和B变体性能相似
  • 在小规模剪枝数据(<1B tokens)中,EfficientLLM-B显著提升准确率

指令微调结果

使用Alpaca数据集(52K指令)微调3个epoch,EfficientLLM-1.1B在Alpaca-Eval上显著超越基线:

  • OLMo-1B
  • ShearedLlama-1.3B
  • TinyLlama-1.1B
  • Llama3.2-1B

七、总结

核心贡献

  1. 提出EfficientLLM系列模型:在100M-1B参数规模实现SoTA性能,超越传统LLM缩放定律
  2. 提出剪枝感知预训练范式:将LLM压缩从后训练提升到预训练阶段,通过数据扩展实现显著性能提升
  3. 探索自动设计架构:首次在现代预训练中实现与人类最佳实践竞争的自动架构设计
  4. 数据高效预训练:50B tokens超越17T tokens训练的Qwen2.5-0.5B,4.15%准确率提升
  5. 高比例剪枝能力:实现70%+剪枝比例且保持合理性能

技术影响

  1. 范式转变:从”训练后压缩”到”训练时压缩”的范式转变
  2. 边缘AI民主化:为移动端和边缘设备提供高性能小模型
  3. 缩放定律突破:证明通过压缩可以超越传统参数缩放定律
  4. 架构搜索自动化:减少对人工架构设计的依赖

局限性

  1. 代码未完全开源:预训练代码和评估代码尚未发布
  2. 源模型依赖:需要预训练好的大模型作为起点
  3. 计算资源需求:需要32-64 A800 GPUs进行训练
  4. 任务范围:主要验证在常识推理任务上,代码/数学等高级能力有待验证
  5. 规模限制:目前验证到1B参数,更大规模的可行性未验证

未来方向

  1. 高级能力保留:探索代码、数学、长上下文等能力的保留
  2. 更大规模扩展:将方法扩展到更大模型和更多参数规模
  3. 多模态扩展:将剪枝感知预训练应用于视觉语言模型
  4. 硬件感知优化:针对特定硬件架构的剪枝策略优化

八、参考资源

论文链接

代码和模型

HuggingFace模型

相关工作

  • MobileLLM: 边缘语言模型的人工架构搜索
  • SmolLM: 高效小语言模型系列
  • LLM-Pruner: LLM结构化剪枝方法
  • SparseGPT: 大语言模型的稀疏化方法
  • ShearedLlama: 从大模型剪枝初始化小模型
  • Qwen2.5-0.5B: 阿里巴巴的小规模语言模型

关键概念

  • Taylor展开剪枝:使用二阶Taylor展开评估参数重要性
  • 双层优化:同时优化剪枝决策和模型权重
  • 显著性驱动:基于梯度和Hessian信息评估参数重要性
  • 结构化剪枝:移除整个神经元/通道/头,保持硬件友好性