面向扩散语言模型的汇聚点感知剪枝

Sink-Aware Pruning for Diffusion Language Models

arXiv: 2602.17664v1

论文信息

标题: Sink-Aware Pruning for Diffusion Language Models

作者: Aidar Myrzakhan, Tianyi Li, Bowei Guo, et al.

发布日期: 2026-02-19

arXiv ID: 2602.17664v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:扩散语言模型(DLM)因迭代去噪而推理成本高昂,现有剪枝方法照搬自回归(AR)模型的 “保留注意力汇聚点” 经验,但该经验是否适用于 DLM 尚未被验证。
  • 核心方法:通过计算汇聚点位置在去噪时间步上的方差(时间方差),发现 DLM 的汇聚点高度不稳定,并提出汇聚感知剪枝——自动识别并抑制这些瞬态汇聚点以优化剪枝决策。
  • 关键结果:汇聚感知剪枝在多个 DLM 上一致超越 Wanda 和 SparseGPT 基线,尤其在 50%-75% 高稀疏度下优势明显(见论文表 1-3)。
  • 主要局限:汇聚统计依赖固定校准集,分布偏移可能降低可靠性;未结合剪枝后微调;在多模态和长上下文场景的验证有限。
  • 适合读者:从事扩散语言模型、大模型压缩、推理加速研究的工程师和研究人员,以及对注意力机制行为分析感兴趣的读者。

论文背景和研究动机

扩散语言模型(DLM)将文本生成建模为迭代去噪过程:从纯噪声开始,逐步去噪直至输出清晰文本。这与自回归(AR)模型逐 Token 生成的方式截然不同,DLM 每次去噪都更新整个 Token 序列。虽然 DLM 在生成质量上展现出竞争力,但其迭代推理大幅增加了计算和内存开销,使得模型压缩(尤其是剪枝)成为实际部署的关键。

现有的 LLM 剪枝方法(如 Wanda、SparseGPT)几乎都继承自 AR 模型,隐含着一个关键假设:注意力汇聚点是不可删减的重要结构。所谓注意力汇聚点,是指序列中某些位置(通常是开头的 BOS Token 或系统提示)持续地、不成比例地吸引大量注意力质量。在 AR 模型中,这些汇聚点充当稳定的全局锚点,帮助传播条件信息并稳定残差流动态。因此,AR 剪枝方法通常会显式保护这些汇聚点,避免灾难性的性能下降。

然而,论文作者质疑这一假设在 DLM 中的有效性。由于 DLM 在每个去噪时间步更新全部 Token,其注意力组织在整个去噪轨迹中持续演化:早期步骤在高噪声下解析全局结构,后期步骤在低噪声下精炼局部语义。这意味着 DLM 的汇聚点可能并非固定不变,而是在去噪过程中漂移、消失或被其他位置取代。

为了量化这一现象,论文定义了两种方差指标:空间方差(各位置在整个轨迹中的平均注意力质量的不均匀程度)和时间方差(汇聚点重心位置随时间步的漂移程度)。分析结果揭示了明确的分歧:AR 模型虽然空间方差高(注意力高度集中于少数位置),但时间方差接近零,汇聚位置极其稳定;而 DLM 相反,空间方差较低但时间方差高出数个数量级,汇聚点位置在去噪过程中持续漂移(见图 4)。这表明,AR 模型中汇聚点作为稳定锚点的特性,在 DLM 中并不成立——很多汇聚点是瞬态的,只在某些噪音阶段发挥作用。

这一观察直接挑战了 AR 剪枝中被奉为圭臬的 “始终保留汇聚点” 原则。

核心方法和技术细节

汇聚感知剪枝方法的核心思路是:先识别哪些 Token 位置是瞬态的不稳定汇聚点,然后在剪枝重要性评估中有意削弱它们对计算的贡献。

汇聚点识别。 对于每个去噪时间步 tt,论文先计算每个 Token 位置 jj 接收到的注意力质量 mt(j)m_t(j),即所有层、所有头中所有查询 token 对位置 jj 的注意力权重之和(公式 9)。若某个位置的注意力质量显著超过其他位置的平均值,即满足 mt(j)>1S−1∑k≠jmt(k)+ϵm_t(j) > \frac{1}{S-1}\sum_{k\neq j} m_t(k) + \epsilon,则被视为当前时间步的汇聚 Token。为获得平滑可微的得分,论文将这一准则通过 Sigmoid 函数松弛为软掩码得分 ϕt(j)\phi_t(j)(公式 11)。

软掩码生成。 在均匀采样的校准时间步集 T\mathcal{T} 上,对上述软得分取平均得到时间步无关的汇聚分数 ϕˉ(j)\bar{\phi}(j)(公式 12)。这个分数越高,表明该位置在整个去噪过程中越频繁地充当汇聚点。若有位置始终是汇聚点(类似于 AR 中的稳定汇聚),其分数会趋近于 1。然后,论文定义位置抑制权重 ωj=1−ϕˉ(j)\omega_j = 1 - \bar{\phi}(j):对于稳定汇聚点,ωj\omega_j 较小,将其注意力贡献大幅衰减;对于普通 Token,ωj\omega_j 接近 1,保持原有贡献。

剪枝重要性重加权。 将抑制后的激活矩阵 X~j,:=ωj⋅Xj,:\widetilde{\mathbf{X}}_{j,:} = \omega_j \cdot \mathbf{X}_{j,:} 代入现有剪枝框架。对于 Wanda,直接使用抑制激活的 L2 范数调整重要性分数;对于 SparseGPT,则用抑制激活计算新 Hessian 矩阵 H~\widetilde{H},后续剪枝和残差重建方程不变。从本质上看,这相当于告诉剪枝器:不要把计算资源分配给那些只在某些噪音阶段短暂成为汇聚点的 Token 位置。

值得注意的是,该方法完全不需要微调或重训练,纯后训练即可插入到 Wanda 或 SparseGPT 的流程中。

创新点和贡献

论文的主要贡献可以归纳为三个层面。

在理论分析层面,论文首次系统揭示了 DLM 中注意力汇聚点的动态特性。此前研究仅观察到了汇聚现象本身,但作者通过精确的时间方差与空间方差分解,证明了 DLM 的汇聚行为与 AR 模型存在本质差异:AR 汇聚点时空聚焦且稳定,DLM 汇聚点空间扩散且瞬态(图 4)。这一发现挑战了跨模型范式直接迁移剪枝经验的常见做法,为设计扩散特化的压缩策略提供了理论依据。

在方法设计层面,汇聚感知剪枝的优雅之处在于其通用性:它不是发明一个全新的剪枝准则,而是对现有准则(Wanda 或 SparseGPT)进行了一次轻量级的、对扩散生成动态敏感的输入变换。这种模块化设计意味着它可以无缝兼容未来更先进的剪枝方法。方法也不需要任何手动指定的硬编码汇聚位置,而完全由校准数据集上的注意力统计自动决定。

在实证效果层面,汇聚感知剪枝在 LLaDA、Dream、LLaDA-1.5 和 MMaDA 四个 DLM 上进行了系统实验。结果表明它一致地匹配或超越了 Wanda 和 SparseGPT 基线(见表 1-5)。在 50% 稀疏度的 LLaDA 上,汇聚感知 SparseGPT 比原始 SparseGPT 平均得分高约 0.45 个百分点;在 Dream 75% 极高稀疏度下,汇聚感知变体仍能保留更多有效模型容量(见表 2)。更重要的是,在需要移除整个注意力头或层的结构化剪枝任务中,汇聚感知剪枝的优势更加明显(见表 4),因为此时每一次剪枝决策的代价都更加高昂。

实践建议

对于希望将汇聚感知剪枝应用到实际 DLM 部署中的从业者,以下建议可能具有参考价值。

剪枝策略选择。 如果压缩目标相对保守(25% 稀疏度以内),标准 Wanda 或 SparseGPT 通常已足够。当需要 50% 或更高的稀疏度时,汇聚感知变体的优势开始明显体现(如图 6 所示,在 LLaDA-1.5 上使用 Wanda 时,75% 稀疏度下汇聚感知剪枝提升近 2 个百分点平均精度)。对延迟敏感但对模型容量有一定容忍度的真实场景,建议从 50% 稀疏度开始尝试。

校准数据配置。 汇聚分数的质量高度依赖校准集对真实推理分布的覆盖。论文使用 WikiText-2 作为校准源,共 128 条截断到 2048 长度。若部署场景涉及特定领域(如代码、法律、医疗),建议构建领域内校准集,否则汇聚统计可能与实际推理时的注意力动态不一致。此外,采样时间步的数量和分布也值得调整,论文未报告不同采样策略下的敏感性。

与量化的联合使用。 论文在局限性部分提到,未来可探索汇聚感知评分与轻量级剪枝后微调以及量化的联合优化。从实践角度看,先执行汇聚感知剪枝再应用量化(如 GPTQ 或 AWQ)是相对安全的路径,因为剪枝后的模型权重已更紧凑,量化误差可能更可控。

计算开销。 方法额外引入的成本主要在校准阶段:需要在多个时间步上提取注意力图。这部分开销在单次后训练压缩中相对模型训练可以忽略不计,但如果需要频繁根据不同分布重建剪枝掩码,累积成本可能变得显著。生产系统中可考虑缓存中间统计量以减少重复计算。