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 |
| HuggingFace | https://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)在大规模参数下实现智能涌现。然而,对于边缘设备部署,存在以下关键挑战:
- 云端成本高:典型LLM(7B-1T参数)需要云端部署和持续网络连接
- 延迟问题:移动边缘应用对实时性要求高
- 隐私顾虑:数据需要上传到云端处理
现有方法的局限:
- 直接预训练:受限于缩放定律,小模型的数据效率低下,且智能涌现仅在较大模型规模出现
- 后训练剪枝:仅使用小型校准数据集,导致显著性能退化
- 知识蒸馏:教师模型(通常7B)的计算量是边缘模型的50倍以上
解决方案概述
本文提出剪枝感知预训练(Pruning-Aware Pretraining),核心思想是在预训练阶段持续进行结构化剪枝,而非传统的后训练剪枝。这具有两个关键特性:
- 数据可扩展(Data-scalable):在LLM中引入最小参数组,持续优化结构化剪枝,将LLM-Pruner和SparseGPT等后训练剪枝方法扩展到预训练阶段
- 架构无关(Architecture-agnostic):使用显著性驱动的剪枝自动设计LLM架构,首次在现代预训练中超越人类设计的最优架构
与现有方法的对比:
- 相比直接预训练:利用更大优化模型的性能,小模型永远无法通过单独预训练达到
- 相比后训练剪枝:用预训练数据扩展剪枝阶段,显著提升压缩质量
三、技术架构
整体框架
剪枝感知预训练被形式化为一个双层优化问题:
外层优化:选择最优的剪枝mini-group g*
内层优化:更新模型权重 w*
训练循环:
1. 梯度下降步骤(权重更新)
2. mini-group优化步骤(剪枝决策)
3. 二阶权重更新(补偿剪枝误差)
核心公式
问题形式化
给定一个优化的大模型 M,剪枝后的模型 M* 可以表示为:
其中 是第t步剪枝的mini-group参数, 是由mini-groups构成的剪枝空间。
递归剪枝
剪枝过程被解耦为t步,可以近似顺序求解:
双层优化
将问题转化为mini-groups g和权重w的双层优化:
外层优化通过Eq.3求解,内层优化通过梯度下降直接求解。梯度下降和mini-group优化交替进行,称为剪枝感知预训练 x1。
二阶Taylor展开
对于优化后的模型,任何权重 w 的损失可以用二阶Taylor展开近似:
其中 , 是Hessian矩阵。
显著性计算
最优mini-group的选择通过显著性评估:
其中三种剪枝类型的显著性分别为:
- 类型I (注意力头剪枝): - 按行求和的注意力输出投影显著性
- 类型II (FFN通道剪枝): - 按行求和的down投影显著性
- 类型III (Stem通道剪枝): - 按列求和的输出层组显著性
二阶权重更新
剪枝后剩余权重的更新公式:
为高效计算Hessian逆,通过求解线性方程:
模型组件
最小剪枝组(Minimal Pruning Groups)
定义三种基本剪枝类型:
| 类型 | 名称 | 剪枝单元 | 耦合参数 |
|---|---|---|---|
| 类型I | 注意力头剪枝 | 每个attention head | Q/K/V输入通道 + O输出通道 |
| 类型II | FFN通道剪枝 | 每个FFN中间通道 | Up/Gate输入通道 + Down输出通道 |
| 类型III | Stem通道剪枝 | 每个transformer stem通道 | Embedding + 所有层Q/K/V/O + FFN + LM Head |
公式表示:
类型I - 注意力mini-group:
类型II - FFN mini-group:
类型III - Stem mini-group:
剪枝空间
原始剪枝空间:
通过mini-group耦合后,每步选择空间降至3,总剪枝空间为 。
EfficientLLM变体
| 变体 | 说明 | 特点 |
|---|---|---|
| EfficientLLM-A | 基础版本 | 使用LLM-Pruner近似Eq.9 |
| EfficientLLM-B | 增强版本 | 在A基础上增加二阶权重更新 |
训练流程
训练阶段
-
剪枝感知预训练阶段:
- 持续进行结构化剪枝和权重优化
- 每步:梯度下降 + mini-group选择 + 二阶权重更新
- 目标:从源模型剪枝到目标模型大小
-
继续预训练阶段:
- 达到目标大小后,继续预训练以提升性能
- 50B-500B tokens
训练超参数
| 参数 | 值 |
|---|---|
| 批量大小 | 1M tokens |
| GPU数量 | 32-64 A800 GPUs |
| 模型规模 | 100M - 1.1B parameters |
| 剪枝比例 | 高达70%+ |
各模型训练详情
| 模型 | 源模型 | 剪枝预训练tokens | 继续预训练tokens |
|---|---|---|---|
| EfficientLLM-134M | SmolLM-360M | 50.3B | 500B |
| EfficientLLM-469M | SmolLM-1.7B | 72.1B | 500B |
| EfficientLLM-1.1B | SmolLM-1.7B | 36.7B | 320B |
数据组成
与源模型SmolLM保持相似的数据分布:
- FineWeb-Edu: 220B tokens
- Cosmopedia v2: 28B tokens
- Python-Edu: 4B tokens
- OpenWebMath: 27.5B tokens (随机采样)
Hessian近似策略
- 显著性检测:使用全局对角Hessian矩阵近似(如LLM-Pruner)
- 权重更新:使用逐层近似
这种解耦策略兼顾了全局显著性检测和局部误差最小化。
四、核心创新
| 创新点 | 说明 | 理论/实验依据 |
|---|---|---|
| 剪枝感知预训练 | 首次将LLM压缩提升到预训练阶段,用预训练数据扩展剪枝 | 即使使用vanilla LLM-Pruner指标,也能超越SoTA剪枝方法 |
| 架构无关设计 | 使用显著性驱动的剪枝自动设计LLM架构,首次超越人类设计 | 自动发现的架构与人类最佳实践(如MobileLLM deep-and-thin)竞争 |
| 最小参数组 | 定义三种最小剪枝单元,将剪枝空间从指数级降至线性 | 原始空间 降至每步选择3种类型 |
| 高效二阶更新 | 解耦Hessian矩阵在显著性检测和权重更新中的应用 | 显著性用全局对角Hessian,权重更新用逐层近似 |
| 数据高效预训练 | 50B tokens超越17T tokens训练的Qwen2.5-0.5B | EfficientLLM-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-134M | 134M | SmolLM-360M | 500B |
| EfficientLLM-469M | 469M | SmolLM-1.7B | 50B / 500B |
| EfficientLLM-1.1B | 1.1B | SmolLM-1.7B | 50B / 320B |
待完成
- 发布技术报告
- 发布HuggingFace模型
- 评估代码
- 预训练代码
- 演示和应用
六、实验结果
基准测试
评估任务
- 世界知识:MMLU
- 常识推理:ARC-c, ARC-e, BoolQ, HellaSwag, OBQA, PIQA, WinoGrande
主要结果 (Table 2)
100M-200M参数规模
| 模型 | #Tokens | #Params | MMLU | ARC-c | ARC-e | BoolQ | HellaSwag | OBQA | PIQA | WinoGrande | Avg. |
|---|---|---|---|---|---|---|---|---|---|---|---|
| OPT-125M | 180B | 125M | 26.02 | 22.87 | 43.31 | 55.44 | 31.37 | 27.80 | 62.62 | 49.80 | 41.89 |
| GPT-neo-125M | 300B | 125M | 26.89 | 23.29 | 43.22 | 61.77 | 30.49 | 26.00 | 62.62 | 51.93 | 42.76 |
| Pythia-160M | 300B | 162M | 26.43 | 22.27 | 37.84 | 43.33 | 29.97 | 26.40 | 58.87 | 49.96 | 38.38 |
| Memba-130M | 1.2T | 130M | 27.65 | 24.49 | 47.56 | 54.68 | 35.11 | 29.00 | 64.69 | 53.35 | 44.13 |
| MobileLLM-LS-125M | 1T | 125M | - | 28.7 | 45.8 | 60.4 | 39.5 | 41.1 | 65.7 | 52.1 | 47.61 |
| SmolLM-135M | 600B | 135M | 30.05 | 29.35 | 61.32 | 59.85 | 42.67 | 34.40 | 68.55 | 52.96 | 49.87 |
| EfficientLLM-A | 500B | 134M | 30.54 | 30.97 | 62.88 | 60.40 | 43.81 | 33.60 | 68.82 | 53.28 | 50.54 |
300M-500M参数规模
| 模型 | #Tokens | #Params | MMLU | ARC-c | ARC-e | BoolQ | HellaSwag | OBQA | PIQA | WinoGrande | Avg. |
|---|---|---|---|---|---|---|---|---|---|---|---|
| OPT-350M | 180B | 331M | 26.96 | 23.98 | 44.02 | 57.80 | 36.63 | 27.80 | 64.91 | 52.96 | 44.01 |
| BLOOM-560M | 350B | 559M | 27.32 | 24.40 | 46.04 | 44.46 | 36.54 | 28.80 | 62.57 | 53.20 | 42.29 |
| Pythia-410M | 300B | 405M | 29.10 | 24.15 | 51.39 | 59.20 | 40.20 | 29.40 | 66.70 | 53.83 | 46.41 |
| MobileLLM-LS-350M | 1T | 345M | - | 32.5 | 54.4 | 62.8 | 50.6 | 45.8 | 69.8 | 57.2 | 53.30 |
| SmolLM-360M | 600B | 362M | 33.89 | 36.26 | 70.16 | 55.23 | 53.51 | 37.60 | 71.38 | 57.22 | 54.48 |
| Qwen2-0.5B | 15T | 494M | 31.85 | 28.50 | 55.05 | 61.25 | 49.16 | 32.80 | 69.75 | 57.22 | 50.53 |
| Qwen2.5-0.5B | 17T | 494M | 33.37 | 32.17 | 64.44 | 61.99 | 52.09 | 35.20 | 70.29 | 56.20 | 53.20 |
| EfficientLLM-A | 50B | 469M | 33.09 | 35.92 | 70.50 | 59.85 | 53.16 | 35.00 | 72.69 | 56.27 | 54.77 |
| EfficientLLM-A | 500B | 469M | 34.54 | 38.40 | 72.10 | 62.42 | 56.84 | 40.40 | 73.83 | 57.46 | 57.35 |
1B参数规模
| 模型 | #Tokens | #Params | MMLU | ARC-c | ARC-e | BoolQ | HellaSwag | OBQA | PIQA | WinoGrande | Avg. |
|---|---|---|---|---|---|---|---|---|---|---|---|
| OPT-1.3B | 180B | 1.3B | 29.57 | 30.03 | 57.49 | 56.54 | 53.66 | 32.80 | 72.31 | 59.04 | 51.70 |
| GPT-neo-1.3B | 380B | 1.3B | 30.00 | 25.94 | 56.31 | 61.90 | 48.99 | 33.40 | 71.00 | 54.62 | 50.31 |
| BLOOM-1.1B | 350B | 1.1B | 29.16 | 25.77 | 51.73 | 59.51 | 43.11 | 29.60 | 67.30 | 54.62 | 47.38 |
| Pythia-1B | 300B | 1.0B | 30.14 | 26.96 | 56.86 | 60.04 | 47.15 | 31.20 | 70.29 | 52.88 | 49.34 |
| TinyLlama-1.1B | 3T | 1.1B | 32.30 | 30.29 | 60.40 | 56.85 | 59.13 | 35.80 | 73.07 | 59.04 | 53.51 |
| ShearedLlama-1.3B | 50B | 1.3B | 31.51 | 29.44 | 61.07 | 61.83 | 59.33 | 34.40 | 73.94 | 58.01 | 54.00 |
| OLMo-1B | 2T | 1.2B | 32.03 | 30.72 | 63.55 | 61.38 | 62.86 | 36.40 | 75.35 | 59.35 | 55.66 |
| Llama3.2-1B | - | 1.2B | 36.31 | 31.48 | 65.28 | 63.88 | 63.69 | 37.40 | 74.59 | 60.54 | 56.69 |
| EfficientLLM-A | 50B | 1.1B | 36.71 | 40.36 | 73.61 | 62.39 | 60.24 | 40.20 | 75.19 | 61.25 | 59.03 |
| EfficientLLM-A | 320B | 1.1B | 37.71 | 42.24 | 73.48 | 67.09 | 64.09 | 41.80 | 75.41 | 61.17 | 60.75 |
与LLM剪枝方法对比 (Llama-7B)
50%剪枝比例
| 模型 | #Tuning | ARC-c | ARC-e | BoolQ | HellaSwag | OBQA | PIQA | WinoGrande | Avg. |
|---|---|---|---|---|---|---|---|---|---|
| MaP | ✓ | 30.63 | 49.32 | 39.69 | 42.49 | 31.40 | 66.81 | 50.67 | 44.43 |
| MvP | ✓ | 26.79 | 44.07 | 59.94 | 40.98 | 31.80 | 63.06 | 55.64 | 46.04 |
| WANDA | ✓ | 34.20 | 42.68 | 50.90 | 38.12 | 38.78 | 57.38 | 55.98 | 45.43 |
| LLM-Pruner | ✓ | 28.24 | 46.46 | 61.47 | 47.56 | 35.20 | 68.82 | 55.09 | 48.98 |
| LoRAPrune | ✓ | 31.62 | 45.13 | 61.88 | 47.86 | 34.98 | 71.53 | 55.01 | 49.72 |
| LoRAShear | ✓ | 32.26 | 47.68 | 62.12 | 48.01 | 34.61 | 71.80 | 56.29 | 50.40 |
| Compresso | ✓ | 27.82 | 48.82 | 60.09 | 39.31 | 33.40 | 66.70 | 51.93 | 46.87 |
| NutePrune | ✗ | 31.74 | 46.59 | 62.20 | 53.87 | 35.80 | 69.91 | 57.77 | 51.13 |
| NutePrune | ✓ | 32.17 | 51.68 | 62.26 | 55.88 | 34.40 | 71.00 | 57.54 | 52.13 |
| EfficientLLM-A | ✗ | 30.80 | 52.15 | 62.29 | 54.70 | 35.20 | 71.33 | 56.75 | 51.89 |
| EfficientLLM-A | ✓ | 34.04 | 64.81 | 64.83 | 60.12 | 34.60 | 73.88 | 61.48 | 56.25 |
70%剪枝比例
| 模型 | #Tuning | ARC-c | ARC-e | BoolQ | HellaSwag | OBQA | PIQA | WinoGrande | Avg. |
|---|---|---|---|---|---|---|---|---|---|
| LLM-Pruner | ✓ | 24.83 | 39.56 | 47.28 | 31.66 | 28.80 | 60.83 | 50.75 | 40.53 |
| NutePrune | ✓ | 26.19 | 42.17 | 62.08 | 39.43 | 30.20 | 62.30 | 51.46 | 44.83 |
| EfficientLLM-A | ✗ | 27.73 | 54.50 | 47.89 | 47.77 | 31.00 | 68.17 | 55.17 | 47.46 |
| EfficientLLM-A | ✓ | 29.95 | 58.59 | 58.13 | 52.02 | 34.60 | 70.08 | 55.96 | 51.33 |
与ShearedLlama对比 (Llama2-7B)
| 模型 | #Pruning | #Tuning | ARC-c | ARC-e | BoolQ | HellaSwag | OBQA | PIQA | WinoGrande | Avg. |
|---|---|---|---|---|---|---|---|---|---|---|
| ShearedLlama-2.7B | 0.4B | – | 26.37 | 49.62 | 59.02 | 47.15 | 33.00 | 66.59 | 53.83 | 47.94 |
| EfficientLLM-A-2.7B | 0.66B | – | 26.71 | 55.51 | 57.31 | 47.58 | 32.20 | 67.14 | 56.51 | 48.99 |
| EfficientLLM-A-2.7B | 10.56B | – | 29.61 | 61.15 | 62.23 | 54.91 | 34.40 | 71.11 | 57.46 | 52.98 |
| ShearedLlama-1.3B | 0.4B | – | 22.78 | 41.08 | 60.18 | 34.66 | 28.20 | 63.00 | 50.67 | 42.94 |
| EfficientLLM-B-1.3B | 0.95B | – | 24.57 | 45.75 | 60.40 | 38.24 | 30.40 | 63.11 | 51.54 | 44.86 |
| ShearedLlama-1.3B | 0.4B | 5.6B | 26.54 | 55.60 | 60.12 | 48.28 | 31.60 | 68.77 | 56.20 | 49.59 |
| EfficientLLM-B-1.3B | 10.56B | 5B | 30.20 | 60.02 | 60.18 | 53.22 | 31.60 | 70.40 | 58.01 | 51.95 |
消融实验
可扩展性分析
设置剪枝步骤与梯度下降步骤的比例为4:1, 2:1, 1:1, 1:9:
| 比例 | 达到目标模型大小所需tokens |
|---|---|
| 4:1 | 2.5B |
| 2:1 | 4.5B |
| 1:1 | 8.4B |
| 1:9 | 72.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
七、总结
核心贡献
- 提出EfficientLLM系列模型:在100M-1B参数规模实现SoTA性能,超越传统LLM缩放定律
- 提出剪枝感知预训练范式:将LLM压缩从后训练提升到预训练阶段,通过数据扩展实现显著性能提升
- 探索自动设计架构:首次在现代预训练中实现与人类最佳实践竞争的自动架构设计
- 数据高效预训练:50B tokens超越17T tokens训练的Qwen2.5-0.5B,4.15%准确率提升
- 高比例剪枝能力:实现70%+剪枝比例且保持合理性能
技术影响
- 范式转变:从”训练后压缩”到”训练时压缩”的范式转变
- 边缘AI民主化:为移动端和边缘设备提供高性能小模型
- 缩放定律突破:证明通过压缩可以超越传统参数缩放定律
- 架构搜索自动化:减少对人工架构设计的依赖
局限性
- 代码未完全开源:预训练代码和评估代码尚未发布
- 源模型依赖:需要预训练好的大模型作为起点
- 计算资源需求:需要32-64 A800 GPUs进行训练
- 任务范围:主要验证在常识推理任务上,代码/数学等高级能力有待验证
- 规模限制:目前验证到1B参数,更大规模的可行性未验证
未来方向
- 高级能力保留:探索代码、数学、长上下文等能力的保留
- 更大规模扩展:将方法扩展到更大模型和更多参数规模
- 多模态扩展:将剪枝感知预训练应用于视觉语言模型
- 硬件感知优化:针对特定硬件架构的剪枝策略优化
八、参考资源
论文链接
- arXiv: https://arxiv.org/abs/2502.06663
- PDF: https://arxiv.org/pdf/2502.06663
- HTML: https://arxiv.org/html/2502.06663v1
代码和模型
- GitHub: https://github.com/Xingrun-Xing2/EfficientLLM
- HuggingFace Collection: https://huggingface.co/collections/xrxing/efficientllm-pruning-aware-pretraining-67a8ecc6a49580b647a6184f
HuggingFace模型
- EfficientLLM-134M: https://huggingface.co/xrxing/EfficientLLM-134M
- EfficientLLM-469M: https://huggingface.co/xrxing/EfficientLLM-469M
- EfficientLLM-1.1B: https://huggingface.co/xrxing/EfficientLLM-1.1B
相关工作
- MobileLLM: 边缘语言模型的人工架构搜索
- SmolLM: 高效小语言模型系列
- LLM-Pruner: LLM结构化剪枝方法
- SparseGPT: 大语言模型的稀疏化方法
- ShearedLlama: 从大模型剪枝初始化小模型
- Qwen2.5-0.5B: 阿里巴巴的小规模语言模型
关键概念
- Taylor展开剪枝:使用二阶Taylor展开评估参数重要性
- 双层优化:同时优化剪枝决策和模型权重
- 显著性驱动:基于梯度和Hessian信息评估参数重要性
- 结构化剪枝:移除整个神经元/通道/头,保持硬件友好性