Back to blog

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 提出了一种新颖的框架,将全局剪枝过程重新定义为可管理的、协调的子问题:

  1. 模块化建模:将 LLM 概念化为模块函数链,一个模块的输出是下一个模块的输入
  2. 辅助变量分解:利用辅助变量将全局剪枝目标分解为等价形式
  3. 交替优化:通过交替优化算法高效求解子问题,每个子问题都有闭式解
  4. 全局最优性:在资源高效优化的同时保持全局最优性

关键创新:通过引入集合 Ω 来控制剪枝的”全局程度”:

  • 当 Ω = {1,2,…,L-1} 时,等价于全局剪枝
  • 当 Ω = ∅ 时,简化为局部剪枝
  • 通过调整 Ω,可以在全局和局部视角之间无缝过渡

三、技术架构

整体框架

SparseLLM 的架构设计基于以下关键观察:

LLM 结构:
┌─────────────────────────────────────────────────────────────┐
│  Decoder Layer ℓ                                            │
│  ┌─────────────┐    ┌─────────────┐    ┌─────────────┐     │
│  │   FFN 模块   │    │  MHA 模块   │    │ Layer Norm  │     │
│  │  (全局剪枝)  │    │  (局部剪枝)  │    │             │     │
│  └─────────────┘    └─────────────┘    └─────────────┘     │
│         ↑                   ↑                   ↑           │
│    权重 W_ℓ            权重 W_MHA           参数 γ, β       │
└─────────────────────────────────────────────────────────────┘

设计策略:

  • FFN 模块:占总参数 2/3 以上,采用全局剪枝
  • MHA 模块:采用局部剪枝(与现有方法一致)
  • 权衡:在计算可行性和剪枝效果之间取得平衡

核心公式

1. 全局剪枝问题 (Eq. 1)

min⁡M,W^L(f(X;M⊙W^),f(X;W))\min_{\mathbf{M},\widehat{\mathbf{W}}} \mathcal{L}(f(\mathbf{X}; \mathbf{M} \odot \widehat{\mathbf{W}}), f(\mathbf{X}; \mathbf{W}))

其中 ⊙ 表示逐元素乘法,𝐌 是稀疏掩码,𝐖̂ 是更新后的权重。

2. 局部剪枝问题 (Eq. 2)

min⁡Mℓ,W^ℓ∥Wℓ⋅Xℓ−(Mℓ⊙W^ℓ)⋅Xℓ∥22\min_{\mathbf{M}_\ell, \widehat{\mathbf{W}}_\ell} \|\mathbf{W}_\ell \cdot \mathbf{X}_\ell - (\mathbf{M}_\ell \odot \widehat{\mathbf{W}}_\ell) \cdot \mathbf{X}_\ell\|_2^2

3. SparseLLM 统一公式 (Eq. 3)

min⁡{W^ℓ},{Mℓ},{aℓ},{zℓ}L(zL,y)\min_{\{\widehat{\mathbf{W}}_\ell\}, \{\mathbf{M}_\ell\}, \{\boldsymbol{a}_\ell\}, \{\mathbf{z}_\ell\}} \mathcal{L}(\mathbf{z}_L, \mathbf{y})

约束条件: zℓ=(Mℓ⊙W^ℓ)aℓ−1,∀ℓ∈[L]\mathbf{z}_\ell = (\mathbf{M}_\ell \odot \widehat{\mathbf{W}}_\ell) \boldsymbol{a}_{\ell-1}, \quad \forall \ell \in [L] aℓ=ϕℓ(zℓ),∀ℓ∈Ω\boldsymbol{a}_\ell = \phi_\ell(\mathbf{z}_\ell), \quad \forall \ell \in \Omega aℓ,zℓ=aℓpre,zℓpre,∀ℓ∈[L−1]∖Ω\boldsymbol{a}_\ell, \mathbf{z}_\ell = \boldsymbol{a}_\ell^{pre}, \mathbf{z}_\ell^{pre}, \quad \forall \ell \in [L-1] \setminus \Omega

关键变量:

  • zℓ\mathbf{z}_\ell:层 ℓ 的输出(辅助变量)
  • aℓ\boldsymbol{a}_\ell:激活值
  • ϕℓ\phi_\ell:非参数化层(激活函数、自注意力、Layer Norm 等)
  • Ω\Omega:参与剪枝的层索引集合

4. 松弛目标函数 (Eq. 4)

L(zL,y)+α∑ℓ∈[L]∥zℓ−(Mℓ⊙W^ℓ)aℓ−1∥22+β∑ℓ∈ΩFFN∥aℓ−ϕℓ(zℓ)∥22\mathcal{L}(\mathbf{z}_L, \mathbf{y}) + \alpha \sum_{\ell \in [L]} \|\mathbf{z}_\ell - (\mathbf{M}_\ell \odot \widehat{\mathbf{W}}_\ell) \boldsymbol{a}_{\ell-1}\|_2^2 + \beta \sum_{\ell \in \Omega_{FFN}} \|\boldsymbol{a}_\ell - \phi_\ell(\mathbf{z}_\ell)\|_2^2

其中 α, β 是控制约束权重的超参数。

5. OPT 模型的子问题求解

权重剪枝:分解 zℓ=Wℓaℓ−1\mathbf{z}_\ell = \mathbf{W}_\ell \boldsymbol{a}_{\ell-1},其中 Wℓ=zℓaℓ−1†\mathbf{W}_\ell = \mathbf{z}_\ell \boldsymbol{a}_{\ell-1}^\dagger(伪逆)

激活值更新 (Eq. 6): (αWℓ+1⊤Wℓ+1+βI)−1(αWℓ+1⊤zℓ+1pre+β⋅ReLU(zℓ))(\alpha \mathbf{W}_{\ell+1}^\top \mathbf{W}_{\ell+1} + \beta \mathbf{I})^{-1} (\alpha \mathbf{W}_{\ell+1}^\top \mathbf{z}_{\ell+1}^{pre} + \beta \cdot \text{ReLU}(\mathbf{z}_\ell))

输出更新 (Eq. 8): zℓ(1)=(Mℓ⊙W^ℓ)aℓ−1pre\mathbf{z}_\ell^{(1)} = (\mathbf{M}_\ell \odot \widehat{\mathbf{W}}_\ell) \boldsymbol{a}_{\ell-1}^{pre} zℓ(2)=(α+β)−1⋅(βaℓ+αzℓ(1))\mathbf{z}_\ell^{(2)} = (\alpha + \beta)^{-1} \cdot (\beta \boldsymbol{a}_\ell + \alpha \mathbf{z}_\ell^{(1)})

6. LLaMA 模型的子问题求解

LLaMA 使用 SiLU 激活函数和门控投影层,解的形式略有不同:

激活值更新 (Eq. 10): (αWℓ+1⊤Wℓ+1+βI)−1(αWℓ+1⊤zℓ+1pre+β⋅SiLU(sℓ)⊙zℓ)(\alpha \mathbf{W}_{\ell+1}^\top \mathbf{W}_{\ell+1} + \beta \mathbf{I})^{-1} (\alpha \mathbf{W}_{\ell+1}^\top \mathbf{z}_{\ell+1}^{pre} + \beta \cdot \text{SiLU}(\mathbf{s}_\ell) \odot \mathbf{z}_\ell)

输出更新 (Eq. 12): zℓ∗=(Mℓ⊙W^ℓ)aℓ−1pre+SiLU(sℓ)⊙aℓSiLU(sℓ)2+1\mathbf{z}_\ell^* = \frac{(\mathbf{M}_\ell \odot \widehat{\mathbf{W}}_\ell) \boldsymbol{a}_{\ell-1}^{pre} + \text{SiLU}(\mathbf{s}_\ell) \odot \boldsymbol{a}_\ell}{\text{SiLU}(\mathbf{s}_\ell)^2 + \mathbf{1}}

时间复杂度分析

SparseLLM 的整体时间复杂度为 O(h³),与 SparseGPT 的每轮复杂度相当:

  1. 权重剪枝步骤:O(nh²)(伪逆计算)+ O(h³)(SparseGPT 求解器)
  2. 激活值更新步骤:O(h³)(矩阵求逆)
  3. 输出更新步骤:较低复杂度

其中 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% WT280% WT290% WT23:4 WT2
Magnitude6420.809998.718209.13-
Wanda21.56142.205692.65-
SparseGPT18.0469.672596.70252.81
SparseLLM17.8258.921350.31128.83

OPT-66B (基线 Perplexity: WT2=9.34, PTB=13.36, C4=10.99)

方法70% WT280% WT290% WT23:4 WT2
SparseGPT9.4528.277803.106594.37
SparseLLM9.3716.457504.174641.80

LLaMA 模型结果 (Table 2)

LLaMA-2 7B (基线 Perplexity: WT2=5.47, PTB=37.91, C4=7.26)

方法70% WT280% WT290% WT23:4 WT2
SparseGPT15.9853.20344.9768.28
SparseLLM16.1549.96225.2364.17

LLaMA-2 13B (基线 Perplexity: WT2=4.88, PTB=50.94, C4=6.73)

方法70% WT280% WT290% WT23:4 WT2
SparseGPT12.9845.59825.9963.48
SparseLLM12.9536.36646.1553.71

关键发现

  1. 高稀疏度优势显著:在 80%-90% 稀疏度下,SparseLLM 相比 SparseGPT 的 perplexity 降低最为明显
  2. 大模型效果更好:随着模型规模增大,SparseLLM 的优势更加突出
  3. 3:4 半结构化稀疏:SparseLLM 在 3:4 半结构化稀疏下表现优异
  4. 计算开销可控:由于交替优化和闭式解,计算时间与 SparseGPT 相当
  5. 内存效率:通过子问题分解,避免了全局剪枝的内存瓶颈

消融实验

  • 校准样本数:在 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 个

注意事项

  1. OPT-350M 不支持:由于潜在的数值稳定性问题
  2. 运行时间:由于迭代交替优化,运行时间约为单次剪枝方法(如 SparseGPT、Wanda)的迭代次数倍
  3. 内存开销:辅助变量引入额外内存开销,对于 LLaMA-2-7B 等大模型,使用较小的校准数据大小(64 或 32)以在 A100 40GB GPU 上运行

七、总结

核心贡献

  1. 统一剪枝框架:提出了通过集合 Ω 控制剪枝全局程度的统一公式,将全局和局部剪枝作为特例
  2. 高效分解算法:利用辅助变量和模块化建模,将全局剪枝分解为可管理的子问题
  3. 闭式解优化:每个子问题都有闭式解,实现快速收敛和计算效率
  4. 显著性能提升:在高稀疏度(>60%)场景下,perplexity 最多降低约 80%
  5. 框架通用性:可无缝集成到 SparseGPT、Wanda 等现有方法中

技术影响

  • LLM 部署:使高稀疏度的 LLM 压缩成为可能,降低推理成本
  • 剪枝研究:提供了一个统一的理论框架来分析和比较不同剪枝方法
  • 硬件友好:3:4 半结构化稀疏特别适合 GPU 加速

局限性

  1. 计算开销:相比单次剪枝方法,运行时间较长(迭代次数倍)
  2. 内存开销:辅助变量引入额外内存消耗,限制了可处理的模型规模
  3. 数值稳定性:交替优化过程对超参数初始化敏感
  4. 模型支持:目前仅支持 OPT 和 LLaMA 系列,更多模型类型待添加

未来方向

  • 优化 GPU 内存消耗,支持更大模型和数据规模
  • 扩展到更多 LLM 架构(如 Mistral、Qwen 等)
  • 探索与量化的结合
  • 研究更高效的交替优化策略

八、参考资源

论文与代码

引用

@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}
}

相关工作

基线模型

  • 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