SparseLLM: Towards Global Pruning for Pre-trained Language Models
一种将全局剪枝分解为可管理子问题的LLM压缩框架
SparseLLM: Towards Global Pruning for Pre-trained Language Models
一、论文概述
| 项目 | 内容 |
|---|---|
| 标题 | SparseLLM: Towards Global Pruning for Pre-trained Language Models |
| 作者 | Guangji Bai, Yijiang Li, Chen Ling, Kibaek Kim, Liang Zhao |
| 机构 | Emory University (Atlanta, GA, USA); Argonne National Laboratory (Lemont, IL, USA) |
| 论文 | arXiv:2402.17946 |
| 代码 | GitHub: BaiTheBest/SparseLLM |
| 发布 | 2024-02-28 (首次提交), 2024-10-31 (最新版本 v4) |
| 会议 | NeurIPS 2024 |
| 领域 | Computer Science > Computation and Language (cs.CL) |
| 许可 | 未明确指定 |
二、核心思想
问题定义
大型语言模型(LLMs)如 LLaMA 和 GPT 在自然语言处理领域取得了卓越成果,但其巨大的计算需求成为主要障碍。剪枝(Pruning) 作为一种关键的压缩策略,通过引入稀疏性来提升内存和计算效率。
然而,现有剪枝方法面临两难困境:
- 全局剪枝(Global Pruning):理论上最优,但需要将整个模型加载到同一 GPU 中,对于现代十亿级参数的 LLM 来说不可行
- 局部剪枝(Local Pruning):逐层独立压缩,效率高但导致次优解,尤其在高稀疏度(>60%)场景下性能显著下降
核心问题:局部剪枝过度约束了中间层的激活值,使得压缩模型与原始模型的中间表示过度对齐,而非关注最终输出的对齐,从而导致全局次优。
解决方案概述
SparseLLM 提出了一种新颖的框架,将全局剪枝过程重新定义为可管理的、协调的子问题:
- 模块化建模:将 LLM 概念化为模块函数链,一个模块的输出是下一个模块的输入
- 辅助变量分解:利用辅助变量将全局剪枝目标分解为等价形式
- 交替优化:通过交替优化算法高效求解子问题,每个子问题都有闭式解
- 全局最优性:在资源高效优化的同时保持全局最优性
关键创新:通过引入集合 Ω 来控制剪枝的”全局程度”:
- 当 Ω = {1,2,…,L-1} 时,等价于全局剪枝
- 当 Ω = ∅ 时,简化为局部剪枝
- 通过调整 Ω,可以在全局和局部视角之间无缝过渡
三、技术架构
整体框架
SparseLLM 的架构设计基于以下关键观察:
LLM 结构:
┌─────────────────────────────────────────────────────────────┐
│ Decoder Layer ℓ │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐ │
│ │ FFN 模块 │ │ MHA 模块 │ │ Layer Norm │ │
│ │ (全局剪枝) │ │ (局部剪枝) │ │ │ │
│ └─────────────┘ └─────────────┘ └─────────────┘ │
│ ↑ ↑ ↑ │
│ 权重 W_ℓ 权重 W_MHA 参数 γ, β │
└─────────────────────────────────────────────────────────────┘
设计策略:
- FFN 模块:占总参数 2/3 以上,采用全局剪枝
- MHA 模块:采用局部剪枝(与现有方法一致)
- 权衡:在计算可行性和剪枝效果之间取得平衡
核心公式
1. 全局剪枝问题 (Eq. 1)
其中 ⊙ 表示逐元素乘法,𝐌 是稀疏掩码,𝐖̂ 是更新后的权重。
2. 局部剪枝问题 (Eq. 2)
3. SparseLLM 统一公式 (Eq. 3)
约束条件:
关键变量:
- :层 ℓ 的输出(辅助变量)
- :激活值
- :非参数化层(激活函数、自注意力、Layer Norm 等)
- :参与剪枝的层索引集合
4. 松弛目标函数 (Eq. 4)
其中 α, β 是控制约束权重的超参数。
5. OPT 模型的子问题求解
权重剪枝:分解 ,其中 (伪逆)
激活值更新 (Eq. 6):
输出更新 (Eq. 8):
6. LLaMA 模型的子问题求解
LLaMA 使用 SiLU 激活函数和门控投影层,解的形式略有不同:
激活值更新 (Eq. 10):
输出更新 (Eq. 12):
时间复杂度分析
SparseLLM 的整体时间复杂度为 O(h³),与 SparseGPT 的每轮复杂度相当:
- 权重剪枝步骤:O(nh²)(伪逆计算)+ O(h³)(SparseGPT 求解器)
- 激活值更新步骤:O(h³)(矩阵求逆)
- 输出更新步骤:较低复杂度
其中 n 是校准样本数,h 是隐藏维度。
四、核心创新
| 创新点 | 说明 | 理论/实验依据 |
|---|---|---|
| 统一剪枝公式 | 通过集合 Ω 控制剪枝的全局程度,统一了全局和局部剪枝 | Remark 4.1:Ω={1,…,L-1} 为全局剪枝,Ω=∅ 为局部剪枝 |
| 辅助变量分解 | 将 LLM 视为模块函数链,引入辅助变量 𝐳_ℓ 和 𝐚_ℓ 解耦参数层和非参数层 | 使全局剪枝问题可分解为可管理的子问题 |
| 闭式解交替优化 | 每个子问题都有闭式解,实现快速收敛 | 类似 ADMM 的多块优化,保证全局收敛性 |
| FFN 全局剪枝策略 | 优先对 FFN 模块(占参数 2/3+)进行全局剪枝,MHA 保持局部剪枝 | 在计算可行性和剪枝效果之间取得平衡 |
| 高稀疏度优势 | 在稀疏度 >60% 时显著优于现有方法 | 实验显示 perplexity 最多降低约 80% |
| 框架通用性 | 可增强 SparseGPT、Wanda 等现有局部剪枝求解器 | 仅需边际额外计算开销 |
五、实验结果
实验设置
- 实现框架:PyTorch + HuggingFace Transformers
- 硬件:NVIDIA A100 GPU
- 校准数据:128 个 2048-token 片段,从 C4 数据集第一分片随机选取(零样本设置)
- 评估指标:Perplexity(WikiText2、PTB、C4)
- 稀疏度:70%、80%、90%、3:4 半结构化
主要结果
OPT 模型结果 (Table 1)
OPT-1.3B (基线 Perplexity: WT2=14.62, PTB=20.29, C4=16.07)
| 方法 | 70% WT2 | 80% WT2 | 90% WT2 | 3:4 WT2 |
|---|---|---|---|---|
| Magnitude | 6420.80 | 9998.71 | 8209.13 | - |
| Wanda | 21.56 | 142.20 | 5692.65 | - |
| SparseGPT | 18.04 | 69.67 | 2596.70 | 252.81 |
| SparseLLM | 17.82 | 58.92 | 1350.31 | 128.83 |
OPT-66B (基线 Perplexity: WT2=9.34, PTB=13.36, C4=10.99)
| 方法 | 70% WT2 | 80% WT2 | 90% WT2 | 3:4 WT2 |
|---|---|---|---|---|
| SparseGPT | 9.45 | 28.27 | 7803.10 | 6594.37 |
| SparseLLM | 9.37 | 16.45 | 7504.17 | 4641.80 |
LLaMA 模型结果 (Table 2)
LLaMA-2 7B (基线 Perplexity: WT2=5.47, PTB=37.91, C4=7.26)
| 方法 | 70% WT2 | 80% WT2 | 90% WT2 | 3:4 WT2 |
|---|---|---|---|---|
| SparseGPT | 15.98 | 53.20 | 344.97 | 68.28 |
| SparseLLM | 16.15 | 49.96 | 225.23 | 64.17 |
LLaMA-2 13B (基线 Perplexity: WT2=4.88, PTB=50.94, C4=6.73)
| 方法 | 70% WT2 | 80% WT2 | 90% WT2 | 3:4 WT2 |
|---|---|---|---|---|
| SparseGPT | 12.98 | 45.59 | 825.99 | 63.48 |
| SparseLLM | 12.95 | 36.36 | 646.15 | 53.71 |
关键发现
- 高稀疏度优势显著:在 80%-90% 稀疏度下,SparseLLM 相比 SparseGPT 的 perplexity 降低最为明显
- 大模型效果更好:随着模型规模增大,SparseLLM 的优势更加突出
- 3:4 半结构化稀疏:SparseLLM 在 3:4 半结构化稀疏下表现优异
- 计算开销可控:由于交替优化和闭式解,计算时间与 SparseGPT 相当
- 内存效率:通过子问题分解,避免了全局剪枝的内存瓶颈
消融实验
- 校准样本数:在 32-64 个样本之间进行敏感性研究
- 剪枝层数:剪枝前 50% 的 Transformer 解码器层以平衡计算资源和性能
- 超参数选择:α 和 β 的选择对性能有重要影响
六、代码实现分析
项目结构
SparseLLM/
├── opt_main.py # OPT 模型剪枝主入口
├── llama_main.py # LLaMA 模型剪枝主入口
├── scripts/ # 复现实验结果的 bash 脚本
├── requirements.txt # 依赖列表
└── README.md # 项目说明
核心依赖
- Python 3.10.14
- PyTorch 2.4.1 (CUDA 12.4)
- Transformers 4.45.1
- Datasets 3.0.1
- numpy 2.1.1
- pandas 2.2.3
- huggingface_hub 0.25.1
- wandb 0.18.2(实验追踪)
使用示例
OPT 模型剪枝:
python opt_main.py \
--model facebook/opt-125m \
--dataset c4 \
--sparsity 0.7
LLaMA-2 模型剪枝:
python llama_main.py \
--model meta-llama/Llama-2-7b-hf \
--dataset c4 \
--sparsity 0.7
半结构化稀疏(2:4):
python opt_main.py \
--model facebook/opt-125m \
--dataset c4 \
--prunen 2 \
--prunem 4
支持的剪枝方法
- 非结构化稀疏:逐个权重剪枝
- 半结构化 N:M 稀疏:
--sparsity_type 2:4:每 4 个权重剪枝 2 个--sparsity_type 4:8:每 8 个权重剪枝 4 个
注意事项
- OPT-350M 不支持:由于潜在的数值稳定性问题
- 运行时间:由于迭代交替优化,运行时间约为单次剪枝方法(如 SparseGPT、Wanda)的迭代次数倍
- 内存开销:辅助变量引入额外内存开销,对于 LLaMA-2-7B 等大模型,使用较小的校准数据大小(64 或 32)以在 A100 40GB GPU 上运行
七、总结
核心贡献
- 统一剪枝框架:提出了通过集合 Ω 控制剪枝全局程度的统一公式,将全局和局部剪枝作为特例
- 高效分解算法:利用辅助变量和模块化建模,将全局剪枝分解为可管理的子问题
- 闭式解优化:每个子问题都有闭式解,实现快速收敛和计算效率
- 显著性能提升:在高稀疏度(>60%)场景下,perplexity 最多降低约 80%
- 框架通用性:可无缝集成到 SparseGPT、Wanda 等现有方法中
技术影响
- LLM 部署:使高稀疏度的 LLM 压缩成为可能,降低推理成本
- 剪枝研究:提供了一个统一的理论框架来分析和比较不同剪枝方法
- 硬件友好:3:4 半结构化稀疏特别适合 GPU 加速
局限性
- 计算开销:相比单次剪枝方法,运行时间较长(迭代次数倍)
- 内存开销:辅助变量引入额外内存消耗,限制了可处理的模型规模
- 数值稳定性:交替优化过程对超参数初始化敏感
- 模型支持:目前仅支持 OPT 和 LLaMA 系列,更多模型类型待添加
未来方向
- 优化 GPU 内存消耗,支持更大模型和数据规模
- 扩展到更多 LLM 架构(如 Mistral、Qwen 等)
- 探索与量化的结合
- 研究更高效的交替优化策略
八、参考资源
论文与代码
- arXiv 论文: https://arxiv.org/abs/2402.17946
- GitHub 代码: https://github.com/BaiTheBest/SparseLLM
- PDF 下载: https://arxiv.org/pdf/2402.17946
- HTML 版本: https://arxiv.org/html/2402.17946v4
引用
@inproceedings{bai2024sparsellm,
title={SparseLLM: Towards Global Pruning of Pre-trained Language Models},
author={Bai, Guangji and Li, Yijiang and Ling, Chen and Kim, Kibaek and Zhao, Liang},
booktitle={The Thirty-eighth Annual Conference on Neural Information Processing Systems},
year={2024}
}
相关工作
- SparseGPT: https://arxiv.org/abs/2301.00774
- Wanda: https://arxiv.org/abs/2306.11695
- LLM-Pruner: 结构化剪枝框架
- LoSparse: 低秩与稀疏矩阵近似结合
- LLM-Shearing: 结构化剪枝方法
基线模型
- OPT 系列: facebook/opt-125m, opt-350m, opt-1.3b, opt-2.7b, opt-6.7b, opt-13b, opt-30b, opt-66b
- LLaMA-2 系列: meta-llama/Llama-2-7b-hf, Llama-2-13b-hf, Llama-2-70b-hf
- LLaMA-3 系列: Llama-3-8b