DIGS: Differentiable, Incremental, Global, Scalable Pruning for Language Models¶
会议: ECCV 2026
论文: ECCV 原文
领域: 多模态VLM
关键词: 模型剪枝、结构化剪枝、大语言模型、梯度范数控制、长周期训练
一句话总结¶
针对基于对偶上升的 \(L_0\) 结构化剪枝在长周期训练中因乘子累积导致的梯度震荡与崩溃问题,DIGS 提出每步门控梯度范数显式约束与闭式求解,并结合单反向传播引擎 GNCL-Opt,首次将大模型结构化剪枝的数据预算平稳扩展至 100B+ Token 级别,在固定稀疏度下实现了性能随训练数据规模单调提升。
研究背景与动机¶
大语言模型(LLM)与多模态大模型(VLM)参数规模的爆炸式增长带来了高昂的训练与部署成本,结构化通道剪枝作为压缩网络核心计算单元的关键技术备受关注。在现有的剪枝范式中,以 SparseGPT 和 Wanda 为代表的事后(post-hoc)非微分方法高度依赖局部统计量,仅需微量校准数据即可完成通道裁剪,但其性能也迅速饱和,无法从更大规模的数据扩增中持续获益。相比之下,基于可微门控和任务损失反向传播的训练期剪枝方法(如基于硬混凝土松弛的 \(L_0\) 正则化剪枝)具备跨层全局预算感知与渐进式调度能力,理论上随着剪枝阶段 Token 预算的增加,任务梯度能够持续重塑保留通道的权重与门控分配,从而进一步缩小剪枝模型与密集基线之间的精度鸿沟。
然而,当尝试将现有的 \(L_0\) 结构化剪枝推向大规模长周期数据训练时,系统遭遇了根本性的数值不稳定性。传统方法普遍依赖增广拉格朗日乘子法(对偶上升,如 Sheared LLaMA 与 L0-Dual)来强制满足全局目标稀疏度。在长周期训练中,拉格朗日乘子 \(\lambda\) 会随时间持续累积暴涨(甚至突破 600),导致门控参数的更新完全被剪枝约束项主导,梯度的微小噪声被急剧放大,在训练中后期引发门控梯度与任务损失的剧烈振荡,甚至出现严重的性能回退。这一内在瓶颈将剪枝阶段的数据规模限制在数个十亿 Token 以内,使得模型在更多数据面前不升反降。
本文的切入角度是打破对偶上升“通过累积惩罚间接压制梯度”的传统范式,直接对每一步反向传播中的门控参数梯度施加显式范数上界,从约束条件推导出单步更新的闭式对偶系数,并辅以单次反向传播的高效实现。核心 idea:提出 DIGS 框架,通过单步门控梯度的 \(L_2\) 范数显式约束闭式求解拉格朗日乘子增量,构建门控饱和自适应衰减与梯度冲突几何折中机制,并借助 GNCL-Opt 单反向传播引擎,彻底消除乘子爆炸诱发的长周期震荡,实现 Token 预算向 100B+ 级别的平稳扩展与性能单调增益。
方法详解¶
整体框架¶
DIGS 的目标是在保持可微分、渐进式调度与全局预算感知优势的同时,彻底解决长周期训练中因拉格朗日乘子失控导致的梯度放大。整个系统以密集预训练或微调后的模型(如 Qwen2.5VL 的 LLM 骨干)为起点,在结构化单元(注意力头、FFN 中间通道)上挂载硬混凝土随机门控 \(z_i \in \{0, 1\}\),并在训练期间通过连续松弛门控 \(\hat{z}_i \in [0, 1]\) 传递梯度。整个剪枝优化流程包括前向双路损失计算、门控梯度的显式范数约束求解、几何冲突自适应折中、单反向梯度重缩放以及基于幂律定标的全局预算控制。
%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
A["输入:密集骨干网络与多模态校准数据"] --> B["梯度范数约束与闭式对偶自适应<br/>显式限制门控合成梯度 L2 范数并闭式求解 λ"]
B --> C["饱和区自适应目标与几何折中控制<br/>结合门控雅可比缩放与任务-剪枝梯度正交投影"]
C --> D["GNCL-Opt 单反向传播引擎<br/>前向无权双路分离与反向剪枝梯度精准缩放"]
D --> E["数据预算幂律定标法则<br/>建立预算与范数目标幂律并在百亿/千亿级扩展"]
E --> F["输出:保持高精度的高稀疏度结构化紧凑模型"]
关键设计¶
1. 梯度范数约束与闭式对偶自适应:消除乘子爆炸并显式控制门控更新步长
传统对偶上升法通过累积稀疏度违反量更新拉格朗日乘子 \(\lambda\),即 \(\hat{z} \leftarrow \hat{z} - \eta (\nabla_{\hat{z}} L_{\text{distill}} + \lambda \nabla_{\hat{z}} L_{\text{prune}})\),长周期训练必然导致 \(\lambda\) 持续膨胀并放大噪声。DIGS 摒弃了乘子的积分累加机制,转而在每一步参数更新时,显式约束门控合成梯度的 \(L_2\) 范数不超过预设阈值 \(c_g > 0\):
$\(\|g_t + \Delta \lambda g_p\|_2 \le c_g\)$
其中 \(g_t = \nabla_{\hat{z}} L(\lambda)\) 为当前步的门控合成梯度,\(g_p = \nabla_{\hat{z}} L_{\text{prune}}\) 为剪枝损失梯度。当该约束存在可行解时,DIGS 令合成梯度饱和在边界上,对应一元二次方程判别式 \(D = (g_t^\top g_p)^2 - \|g_p\|^2 \|g_t\|^2 + \|g_p\|^2 c_g^2\)。若 \(D \ge 0\),选取非负可行根;若约束无法满足(即 \(D < 0\)),则切换为最小二乘分支以获取最小可达范数:
$\(\Delta \lambda = \begin{cases} \frac{-g_t^\top g_p + \sqrt{D}}{\|g_p\|^2 + \varepsilon_{\text{den}}}, & D \ge 0 \\[6pt] -\frac{g_t^\top g_p}{\|g_p\|^2 + \varepsilon_{\text{den}}}, & D < 0 \end{cases}\)$
最终乘子通过非负截断更新 \(\lambda \leftarrow \max(\lambda + \Delta \lambda, 0)\)。该设计将乘子更新与门控梯度范数直接绑定,乘子数值在经历初期爬升后始终被约束在低位(绝大多数步数小于 100),从数学机制上杜绝了对偶乘子爆炸。
2. 饱和区自适应目标与几何折中控制:缓解门控饱和漂移并平衡任务拟合与稀疏化
硬混凝土门控在训练中后期逐渐趋近两极(0 或 1),其类 Sigmoid 雅可比导数急剧收缩,使得剪枝梯度范数 \(\|g_p\|\) 急剧衰减。若依然维持固定的 \(c_g\),闭式解分母中的极小值将导致 \(\Delta \lambda\) 出现虚假虚高。为此,DIGS 引入平滑饱和自适应缩放机制,将有效范数目标定义为:
$\(c_g^{\text{eff}} = c_g \cdot \frac{\|g_p\|}{\|g_p\| + \varepsilon_{\text{scale}}}\)$
当门控处于活跃区时保持原始控制强度,而在门控高度饱和、剪枝方向已基本确立的收敛末期自动衰减范数目标。与此同时,DIGS 具备精细的几何折中特性:设 \(g_d = \nabla_{\hat{z}} L_{\text{distill}}\),当任务蒸馏梯度与剪枝梯度夹角小于 \(90^\circ\)(方向协同)且 \(\|g_d\|\) 较大时,算法自适应缩小 \(\lambda\),避免过度激进剪枝;当两者夹角大于 \(90^\circ\)(相互拮抗)且 \(\|g_d\|\) 增强时,算法适度放大 \(\lambda\) 抵抗任务梯度的反向压制;而在极度不可行区(\(D < 0\) 且 \(\|g_d\|\) 极大),最小二乘解直接消除合成梯度在 \(g_p\) 上的投影分量,主动暂停剪枝推进,优先保证蒸馏主任务平稳拟合。
3. GNCL-Opt 单反向传播引擎:消除双重反向计算开销的高效系统实现
直接按数学公式计算 DIGS 控制律需要先反向传播一次获取门控梯度统计量以求解 \(\lambda\),再执行第二次反向传播完成全模型参数更新,导致训练开销翻倍。为此,DIGS 设计了梯度范数控制拉格朗日优化引擎(GNCL-Opt)。在前向传播阶段,GNCL-Opt 将剪枝门控参数解耦复制为两个计算分支,分别流向蒸馏损失 \(L_{\text{distill}}\) 和剪枝损失 \(L_{\text{prune}}\),总目标保持无加权的朴素相加 \(L = L_{\text{distill}} + L_{\text{prune}}\)。在反向传播阶段,自定义 Backward 控制钩子同时拦截到两路门控梯度信号,立即原位计算出 \(\lambda\),仅对剪枝分支的梯度施加 \(\lambda\) 倍缩放后直接合并。这一实现与在前向加权 \(L_{\text{distill}} + \lambda L_{\text{prune}}\) 产生完全相同的参数更新量,但将计算流程压缩为单次反向传播,消除了额外反向传播的内存占用与显卡算力开销。
4. 数据预算幂律定标法则:根据目标计算预算精确反解超参数与规模化剪枝
由于 DIGS 的剪枝推进速率并非由人为预设的阶梯 warmup 强行指定,而是由梯度范数上限 \(c_g\) 自然涌现,因此选择不同的 \(c_g\) 会间接决定模型达到目标稀疏度所需的 Token 总消耗。作者通过在固定稀疏度目标 \(R\)(如 30%、45%、60%)下进行数轮短周期的对数网格校准实验,发现剪枝 Token 预算 \(\mathbf{B}\) 与 \(c_g\) 之间严格服从幂律分布:
$\(\mathbf{B}(R) \approx a(R) \cdot c_g^{-m(R)} + b(R)\)$
其中 \(a(R)\) 为尺度系数,\(m(R)\) 为幂指数,\(b(R)\) 为修正偏移量。在已知可用算力或预期训练数据规模 \(\mathbf{B}^\star\) 的情况下,研究人员可直接通过反向公式精确解出最优梯度阈值 \(c_g^\star = \left(\frac{a(R)}{\mathbf{B}^\star - b(R)}\right)^{1/m(R)}\),使大规模长周期剪枝实验摆脱了盲目调参,为扩展至 100B+ Token 级别的稳定剪枝提供了可靠的理论与工程抓手。
损失函数 / 训练策略¶
整个剪枝管线以冻结视觉编码器与多模态投影层的 Qwen2.5VL-3B 为核心,仅针对其 LLM 骨干进行剪枝。蒸馏损失采用纯净的前向 KL 散度(不对输入 Prompt 计算,仅在自回归响应 Token 上监督): $\(L_{\text{distill}} = \mathrm{KL}(p_{\text{teacher}} \,\|\, p_{\text{student}})\)$ 全局加权稀疏度损失定义为参数质量加权的期望移除比例 \(R(\alpha) = \frac{1}{S} \sum_{i=1}^N s_i \pi_i(\alpha_i)\),其中 \(s_i\) 为对应通道移除时节省的实际参数量。模型参数学习率设为 \(1\times 10^{-4}\),硬混凝土门控参数 \(\alpha\) 学习率设为 \(1\times 10^{-2}\)。训练前执行 1000 步密集微调 warmup 以解耦学习率热身与剪枝调度,后续正式剪枝无需任何人为设计的稀疏度 warmup 过程,全程依靠 DIGS 闭式范数控制自主演化。
实验关键数据¶
主实验¶
主实验在多模态权威评测工具包 VLMEvalKit 上展开,涵盖多模态综合理解、多学科推理、视觉数学、图表与文档 OCR 等 8 个主流基准:MMBench、MMStar (MMSt)、MMMU (MMMU-Dev-Val)、MathVista (MathV)、HallusionBench (Hallu)、AI2D、OCRBench (OCR) 与 MMVet。测试统一在 30% 与 60% 目标稀疏度下与强基线 L0-Dual 进行严格对齐比较。
在 30% 稀疏度下,当训练数据从 5.7B 扩展至 24.5B 时,DIGS 与 L0-Dual 呈现出截然相反的演化趋势:
| 方法 (Run / Ckpt) | 消耗 Token (B) | MMBench / Hallu | MMSt / AI2D | MMMU / OCR | MathV / MMVet | 综合平均 (Avg) | 相比密集基线 (∆Acc) |
|---|---|---|---|---|---|---|---|
| Qwen2.5VL-3B (密集基准) | – | 74.75 / 46.60 | 56.30 / 81.40 | 51.20 / 82.80 | 61.20 / 65.90 | 65.01 | 0.00 |
| L0-Dual FR5.7B@30% | 3.1 | 68.17 / 49.55 | 48.46 / 70.17 | 40.33 / 78.20 | 53.70 / 57.52 | 58.26 | -6.75 |
| DIGS FR5.7B@30% | 2.7 | 64.33 / 44.96 | 51.26 / 67.09 | 37.33 / 76.70 | 53.30 / 53.48 | 56.04 | -8.97 |
| L0-Dual FR24.5B@30% | 11.9 | 63.40 / 48.24 | 49.86 / 67.64 | 40.77 / 76.60 | 52.70 / 55.59 | 56.85 | -8.16 |
| DIGS FR24.5B@30% | 10.3 | 68.54 / 44.75 | 52.33 / 70.88 | 41.89 / 78.50 | 57.30 / 55.87 | 58.75 | -6.26 |
在挑战性更高的 60% 稀疏度下(模型骨干超过一半通道被剔除),长周期训练下的稳定性差异被进一步放大:
| 方法 (Run / Ckpt) | 消耗 Token (B) | MMBench / Hallu | MMSt / AI2D | MMMU / OCR | MathV / MMVet | 综合平均 (Avg) | 相比密集基线 (∆Acc) |
|---|---|---|---|---|---|---|---|
| Qwen2.5VL-3B (密集基准) | – | 74.75 / 46.60 | 56.30 / 81.40 | 51.20 / 82.80 | 61.20 / 65.90 | 65.01 | 0.00 |
| L0-Dual FR5.7B@60% | 5.70 | 41.94 / 43.46 | 43.66 / 47.05 | 33.00 / 71.30 | 37.90 / 37.80 | 44.51 | -20.50 |
| DIGS FR5.7B@60% | 5.70 | 50.70 / 36.53 | 44.60 / 55.21 | 35.00 / 71.50 | 47.69 / 42.20 | 47.92 | -17.09 |
| L0-Dual FR24.5B@60% | 24.50 | 26.96 / 40.66 | 38.41 / 41.25 | 30.77 / 63.40 | 36.70 / 33.62 | 38.97 | -26.04 |
| DIGS FR24.5B@60% | 24.50 | 59.74 / 44.43 | 45.80 / 57.41 | 36.88 / 74.90 | 51.40 / 44.91 | 51.93 | -13.08 |
在纯 LLM 领域(Gemma 与 LLaMA2 系列骨干),DIGS 在相同的 SFT-then-prune 协议下对比现有剪枝算法表现卓越:
| 稀疏度 | 剪枝算法 | Gemma-2B | Gemma-7B | LLaMA2-7B | LLaMA2-13B |
|---|---|---|---|---|---|
| 0% (Dense) | 原始模型 | 53.82 | 71.59 | 58.76 | 66.74 |
| 25% | LLM-Pruner | 40.20 | 50.29 | 39.72 | 39.82 |
| 25% | SliceGPT | 39.72 | 52.13 | 41.97 | 46.75 |
| 25% | PAT | 52.98 | 66.68 | 60.02 | 66.58 |
| 25% | DIGS (本文) | 53.65 | 70.73 | 57.45 | 67.12 |
| 30% | LLM-Pruner | 40.05 | 41.35 | 39.73 | 39.70 |
| 30% | SliceGPT | 39.89 | 44.30 | 40.14 | 46.19 |
| 30% | PAT | 45.33 | 64.58 | 57.81 | 65.15 |
| 30% | DIGS (本文) | 49.81 | 70.12 | 58.27 | 66.15 |
消融实验¶
消融实验围绕“常规稳定策略是否足以拯救对偶上升”以及“DIGS 扩展至百亿与千亿 Token 级的表现”展开。在评估计算成本时采用 Avg3 指标(MMStar、MathVista-Mini 与 OCRBench 的均值)。
关于常规稳定器的消融表明,事后启发式截断无法替代自适应对偶控制:
| 变体方案 | 训练 Token 量 | Avg3 综合表现 | 核心机理与现象剖析 |
|---|---|---|---|
| L0-Dual (基线) | ~12B | 46.32 | 未加稳定器,长周期下乘子 \(\lambda\) 突破 600,引发后期剧烈震荡 |
| L0-Dual + 梯度裁剪 (clip) | ~12B | 45.69 | 合成后全局裁剪同时削弱了任务蒸馏和剪枝梯度,精度反而进一步下滑 |
| L0-Dual + 乘子硬截断 (\(\lambda \le 50\)) | ~19B | 52.74 | 人工设上限减缓了发散,但严重拖慢剪枝进度,达到 60% 稀疏度需多耗费 ~7B Token |
| DIGS (本文) | ~12B | 56.29 | 在反向融合前闭式自适应约束门控梯度范数,进度与精度兼优 |
梯度约束系数 \(c_g\) 的敏感度实验展示了平滑的算力换精度规律(60% 目标稀疏度):
| 控制系数 \(c_g\) | 实际消耗 Token | MMStar | MathVista-Mini | OCRBench | Avg3 均值 |
|---|---|---|---|---|---|
| \(1.2 \times 10^{-2}\) | 1.06B | 39.80 | 34.60 | 62.80 | 45.73 |
| \(4.7 \times 10^{-3}\) | 2.42B | 42.33 | 42.30 | 68.60 | 51.08 |
| \(1.9 \times 10^{-3}\) | 5.71B | 44.60 | 47.69 | 71.50 | 54.60 |
| \(7.5 \times 10^{-4}\) | 12.46B | 47.07 | 49.10 | 72.70 | 56.29 |
| \(3.0 \times 10^{-4}\) | 24.50B | 45.80 | 51.40 | 74.90 | 57.37 |
基于上述定标法则,作者在 33% 剪枝比例(将 3B 模型物理削减至 2B)下将训练推进至 100B+ Token 规模,各阶段增益归因如下:
| 训练阶段 | 剪枝 Token 量 | 恢复与微调 Token 量 | 8 个多模态基准综合 Avg | 相对密集原模变化 |
|---|---|---|---|---|
| Dense 原始基线 | – | – | 65.01 | 0.00 |
| Prune-only (仅剪枝阶段) | 114B | – | 61.52 | -3.49 |
| + post-pruning 蒸馏恢复 | 114B | 20B | 63.94 | -1.07 |
| + FineVision SFT 对齐微调 | 114B | 20B + 1 epoch SFT | 64.32 | -0.69 |
关键发现¶
- 数据扩展的单调性逆转:在 60% 极高稀疏度下,L0-Dual 在 5.7B Token 处即达到性能峰值(44.51),数据增加到 24.5B 时暴跌至 38.97(-5.54);而 DIGS 随着数据量增加从 47.92 稳步攀升至 51.93(+4.01),彻底逆转了传统剪枝在大规模数据下的负迁移现象。
- 训练收敛动力学优化:训练监控显示,L0-Dual 的最终蒸馏损失停留在 ~1.47 且伴随巨大方差,乘子 \(\lambda\) 超过 600;DIGS 将乘子平稳维持在 100 以内,最终蒸馏损失大幅降低至 ~0.884,方差极小且平稳收敛。
- 千亿 Token 几乎无损剪枝:通过 114B 剪枝 + 20B 蒸馏 + 轻量 SFT,剪除 33% 通道后的 2B 模型在 8 大多模态榜单上的综合得分为 64.32,距离 3B 密集模型(65.01)仅差距 0.69 分,证明充分的数据规模可完全弥补结构化剪枝带来的容量折损。
亮点与洞察¶
- 从“被动惩罚累积”到“主动梯度限幅”的优化范式革新:传统对偶上升将未满足约束视为误差进行积分累加,不可避免造成梯度的指数级放大;DIGS 将控制点前移到每步更新的反向传播路径上,从几何空间直接限定更新步长,构成了严密的闭环负反馈调节。
- GNCL-Opt 单反向传播的高效系统融合:巧妙通过前向分支解耦与反向原位缩放合并,完美避开了二次反向传播带来的双倍显存与算力开销,使得大模型在超大规模数据下的微积分剪枝在工程上极具可行性。
- 可反解的数据预算定标法则:论文提炼的幂律经验公式 \(\mathbf{B}(R) \approx a(R) c_g^{-m(R)} + b(R)\) 建立了超参数 \(c_g\) 与训练 Token 预算的量化关系,使深度学习系统工程人员能像设定预训练步数一样,精确按预算计划剪枝数据。
局限与展望¶
- 剪枝范围局限于 LLM 语言骨干:本文验证实验中冻结了 ViT 视觉编码器与多模态投影层,尚未探索全模态(视觉与语言联合)的端到端结构化剪枝与跨模态冗余权衡。
- 超参标定需预先进行轻量校准:幂律定标公式中的参数 \(a, m, b\) 依赖于具体网络架构与数据分布,在迁移至全新模型时仍需消耗少量算力运行 4-6 组短周期校准。
- 向动态与细粒度稀疏化拓展:该梯度范数约束优化框架本质上具备通用性,未来有望推广至 Mixture-of-Experts(MoE)专家剪枝、KV Cache 动态上下文剪枝以及半结构化稀疏领域。
相关工作与启发¶
- vs L0-Dual (Sheared LLaMA, Wang et al.):二者均采用硬混凝土门控进行全局结构化剪枝,但 L0-Dual 依赖增广拉格朗日乘子累加,导致长周期严重震荡且 Token 扩展受限于数个 B;DIGS 采用单步梯度范数显式约束与闭式求解,支持 100B+ Token 平稳训练且精度单调提升。
- vs Post-hoc 剪枝 (SparseGPT, Wanda):事后剪枝无需反向传播训练、耗时极短,但精度容易在小数据量后迅速饱和;DIGS 属于深度训练期剪枝,能够持续将千亿 Token 的任务监督信号转化为更精准的通道去留决策。
- vs PAT (Pruning-Aware Tuning):PAT 采用多阶段启发式渐进调优(通常使用约 0.256B Token);DIGS 在统一协议下于 Gemma 和 LLaMA2 上均取得了更优的保留精度,尤其在 30% 较高稀疏度下优势更加显著。
评分¶
- 新颖性: ⭐⭐⭐⭐⭐ 首次揭示了拉格朗日 \(L_0\) 剪枝长周期乘子爆炸的本质成因,提出了简洁优美且具备闭式解的梯度范数控制新范式。
- 实验充分度: ⭐⭐⭐⭐⭐ 涵盖 8 大多模态评测基准与 4 个主流纯语言模型,提供了从数亿到 114B Token 的完备定标与消融实验。
- 写作质量: ⭐⭐⭐⭐⭐ 数学推导严密清晰,几何直觉与系统工程实现(GNCL-Opt)阐述透彻,实验图表信息量极其丰富。
- 价值: ⭐⭐⭐⭐⭐ 为大模型结构化剪枝开辟了与数据规模(Data Scaling)协同提升的新路径,工程落地价值极高。