动态稀疏:从剪枝到训练范式重塑神经网络高效训练
2026/9/23 6:47:39 网站建设 项目流程

1. 一张稀疏的网络,凭什么能比稠密网络训得更好?

先说个反直觉的现象:过去两年我在做推荐系统模型压缩时,发现单纯把一个大模型剪到 10% 密度,精度损失通常在 3 到 5 个点;但用动态稀疏方式从头训练一个 10% 密度的稀疏网络,精度损失往往能压到 1 个点以内,有些任务甚至反超稠密基线。这件事促使我把 Dynamic Sparsity 从“一种模型压缩 trick”重新理解成“一种新的训练范式”。

Dynamic Sparsity,翻译过来叫动态稀疏,核心思想就是网络在训练过程中,稀疏连接的模式不是训练前定死的,而是每隔若干步动态更新一次——去掉一批当前重要性最低的权重,同时补回一批新的权重,让“哪些连接存在”这件事也跟着梯度迭代一起演化。与之对应的静态稀疏,则是训练前固定一个 mask,训练过程中 mask 永远不变,压缩率再高也只是在固定拓扑里调权重。

之所以说它是训练范式而非简单的剪枝技巧,是因为动态稀疏真正改变的,是网络在整个优化轨迹上探索的结构空间。静态稀疏相当于把一个网络塞进一条固定的小巷里,只能在巷子里挪动;动态稀疏则每隔一段时间给你换一条巷子,虽然每条巷子都不宽,但组合起来能覆盖的区域远超单条巷子。大量实验已经表明,动态稀疏训练的稀疏网络,最终收敛点和稠密网络的收敛点在泛化能力上相当接近,甚至因为结构扰动带来的隐式正则化,有时还更好。

这篇文章会把这套东西从原理到工程实现完整拆开。内容包括:动态稀疏最核心的机制拆解、这套机制如何体现在实际训练流程中、SET、RigL、DSR 等经典算法各自做了什么、在大模型时代它又以什么形态回归,以及我在落地过程中踩过的坑和总结出的实操经验。适合正在做模型压缩、边缘端部署、大模型高效训练的同学参考。

2. “剪完再训”和“边剪边训”:动态稀疏到底动在哪里

理解动态稀疏之前,必须先把静态剪枝和动态稀疏的训练流程差异看清楚。很多人以为动态稀疏只是“剪枝频率高一点”,这是最常见的误解。

2.1 静态剪枝的标准流程:train-prune-finetune 三步曲

静态剪枝的主流做法是三段式:先正常训练一个稠密网络到收敛,然后根据某种重要性指标把不重要的连接剪掉,形成稀疏 mask,最后在这个 mask 固定的情况下做若干轮微调。整个过程里,mask 一旦生成就不再变化。

这套流程的优点是简单、稳定、工程上好实现;缺点是稀疏结构一经确定就没有回头路。剪错了就是错了,微调只能修正权重数值,没有办法重新长出那些被剪掉的连接。对于比较大的模型,比如百亿参数的语言模型,这种“先训后剪”的方式还有另一个问题:前期训练成本一分没省,只有推理阶段能拿到稀疏红利。

还有一个更隐蔽的问题:预训练模型里学到的连接重要性,未必是稀疏结构最优解。因为你只能在你已经训练出的那个局部最优附近的权重分布上做重要性判断,等于用“这个区域的先验知识”限制了稀疏搜索空间。如果一开始就用高稀疏度约束训练,网络完全可能找到另一组更优的连接组合。

2.2 动态稀疏的循环结构:剪枝-生长-再训练

动态稀疏的训练流程是一个不断循环的结构,核心有三个阶段,循环往复直到训练结束:

  1. 适度训练:在当前的稀疏 mask 下训练若干步,让现有连接对应的权重充分收敛。
  2. 剪枝:按照重要性指标移除一部分当前“最不重要”的连接,释放出对应的稀疏度预算。
  3. 生长:在释放的预算里,按某种策略选择新的连接位置,把 mask 重新补满到目标稀疏度。

之后带着新的 mask 继续训练,重复这个过程。这整个流程中,mask 不是一成不变的,而是“热点”随时在权重矩阵的不同位置间迁移。这种剪枝和生长的组合拳,在形态上很像进化算法里“选择—变异—保留”的框架,只不过它作用在单次梯度训练过程中,尺度极小、频率极高。

要特别注意“剪枝”和“生长”通常是等额的:每轮剪掉多少连接,就必须长出多少连接,让整体稀疏度保持不变。这样设计是为了稳定训练,因为稀疏度突升突降会让 loss 曲线剧烈震荡。而最终模型使用的就是训练结束时的那个 mask,不需要额外的 prune 步骤。

2.3 三类重要性指标对比:幅度、梯度、动量

决定每轮剪掉哪些权重、长出哪些权重,是动态稀疏的“灵魂”所在。核心在怎么定义“重要性”。最常见的做法有三种:

指标类型计算方式特点典型代表
幅度剪枝按权重绝对值排序,最小的剪掉简单高效,但只反映当前数值大小,不考虑未来潜力SET、SNFS
梯度幅值用权重×梯度的绝对值作为重要性信号能感知当前优化方向,但梯度噪声大,需要平滑RigL
动量累积基于历史梯度EMA与权重的乘积抗噪声强,结构更稳定,是目前综合效果最好的方案RigL的改进版、ADP

幅度剪枝最直观:权重绝对值小,说明对输出的影响弱,剪掉它对 loss 的影响最小。这个是经典剪枝论文里的常规思路,实现起来只有几行代码。但动态稀疏场景下“权重小”不代表“连接没用”——有些连接现在小是因为还没被充分训练,给你机会长得更大。

梯度类的指标则更“激进”:一个权重虽然当前数值小,但如果它的梯度很大,说明 loss 对它的响应很敏感,剪掉它会阻碍训练进程。把权重和梯度结合起来,公式一般是importance = |weight × gradient|,这个值越大表示该权重“既在实际生效,又在往关键方向演化”,应当保留。

动量累积是在梯度类基础上进一步平滑噪声。我实测下来,小 batch size 训练时,原始梯度方差极大,直接用weight × gradient剪枝很容易把本该保留的重要连接误杀。引入梯度动量(其实就是 Adam 里类似m_t那一项)之后,剪枝结果稳定很多,这也是 RigL 后续版本默认用动量而非原始梯度的原因。

3. 动态稀疏的核心引擎:剪枝、生长、稀疏度控制怎么做

上一节是概念层,这一节讲实现层。动态稀疏框架里,剪枝和生长策略是成对出现的,你选了什么样的剪枝标准,就必须配套什么样的生长策略,二者共同决定每一轮 mask 的演化轨迹。

3.1 剪枝策略:全局排序还是分层配额

剪枝时最直接的做法是“全局剪枝”:把所有参数的重要性值放在一起排序,一刀切,剪掉全局最小的那部分。全局剪枝的好处是每一步都在全网络范围内做最优资源配置——哪个层该多剪、哪个层该少剪,不完全由人为指定,而是由数据自己说话。

但全局剪枝有个工程隐患:它可能导致某一层的连接被剪得过于稀疏,甚至剪到只剩个位数连接,训练时这一层直接退化失效,反向传播的梯度链断裂。所以更稳的是“分层配额制”:每一层维持一个目标稀疏度,每轮只在该层内部按重要性排序剪枝。层与层之间的稀疏度分布,则可以通过类似“Erdos-Renyi”分布来初始化。

Erdos-Renyi 分布的规则很简单:某层的可剪参数数量与该层权重矩阵的两个维度相关,公式大致是n = (n_in + n_out) / (n_in × n_out)的倒数。翻译成人话就是:输入输出维度越大的层,保留的连接数占比越高。因为大矩阵天然有更多冗余,但完全按比例剪又会把关键信息都剪没了,ER 分布给出的就是一个经验上更合理的配额方案。

选择分层剪枝还有一层考虑:工程实现上并行度更好。全局剪枝要做一次全参数 sort,在百亿模型上这一步的时间和显存开销都不可忽视;分层剪枝天然可以分设备并行,跟分布式训练无缝衔接。

3.2 生长策略:补回连接的最优方式

剪枝释放了一批连接位置,接下来要让网络长回同样数量的新连接。生长策略决定了新连接长在哪儿。主要有四类:

  • 随机生长:在未连接的候选位置中等概率随机选一批补上。这是 SET、SNFS 等早期算法的做法,简单粗暴但效果意外地好。原因在于,动态稀疏本身就在不断探索拓扑空间,随机生长提供了足够的探索随机性,帮助逃离局部结构极优。

  • 梯度最大生长:计算所有未连接位置的梯度幅值,选择梯度最大的位置补回。逻辑是:梯度大说明这个位置“诱导 loss 下降的欲望”强,在这个位置建立连接能更快降低误差。RigL 初版就是这么做的。

  • 反向稀疏补全:把已剪枝位置和未连接位置统一考虑,按同一个重要性指标排序,剪掉最不重要的,同时把最重要的未连接位置长回来。这其实是把剪和长合并成一次全局排序操作,保证“每剪掉一个,一定长回一个当前最优的”。

  • 周期生长:不是每轮都长,而是每隔 N 步统一生长一轮。这种做法是为了让 mask 变化频率与学习率退火节奏对齐。训练后期网络趋近收敛,频繁改结构会扰动已经学好的特征,周期拉长反而有助于稳定收敛。

3.3 稀疏度控制:训练过程中要不要一直不变

很多人做动态稀疏时,把目标稀疏度设成一个固定值(比如 90%),整个训练过程不变。主流做法确实如此,但更精细的方案是“稀疏度升温”,英文通常叫 Gradual Sparsity Increase。

道理很简单:训练初期网络还在学习基础特征,如果一上来就 90% 稀疏,每轮更新的有效参数量太少,模型可能永远学不会。所以比较稳的做法是:从 0% 或较低的稀疏度开始,随着训练轮数线性或指数提升到目标稀疏度。这个思路类似于学习率 warmup,给网络一个“先学能力、再压缩结构”的缓冲期。

我在多组实验里对比过固定稀疏度和升温稀疏度。在 CIFAR 和 ImageNet 这种标准视觉任务上,升温方案能稳定提升 1 到 2 个点的精度,而且在 95% 以上超低密度区间,升温几乎是必须的,否则会出现严重的梯度消失或训练崩溃。

升温节奏有个经验公式可以参考:sparsity_t = target_sparsity × (1 - (1 - t/T)^power)power取 3 时曲线后段增长最平滑,训练末期 mask 变化幅度小,收敛稳定。如果power=1线性增长,训练后期 mask 变化还是太大,loss 容易在最后阶段翘尾回升。

3.4 DSR 的重参数化技巧:让稀疏 mask 直接参与反向传播

上面讲的剪枝和生长都是在离散 mask 上操作,这个过程不可导。大多数动态稀疏算法都是把 mask 当“开关”用,梯度不经过 mask 本身,这也意味着网络没法通过学习来调整“哪些连接重要”这种高层行为。

Deep Sparse Rewiring(DSR)提出的重参数化思路,解决的就是这个问题。它的做法是:不给每个权重二进制的 0/1 mask,而是给每个权重一个连续分布参数,这个参数控制该连接“存在”的概率。训练时从分布中采样出实际 mask,采样过程用 Gumbel-Softmax 之类的连续近似替代,这样 mask 的“概率参数”就能参与反向传播。

换句话说,普通动态稀疏是“外力”决定谁死谁活,DSR 是让网络自己学习该让谁死谁活。DSR 的核心数学形式类似变分推断——每个连接的重要性被建模成一个可学习的分布,每次前向采样得到的稀疏结构天然带了随机性,相当于在做结构层面的数据增强。

不过说实话,DSR 在中小规模模型上效果不错,但在大规模训练里工程复杂度偏高。一个原因是它对每个权重都要额外维护分布参数,显存开销翻倍;另一个原因是采样带来的随机性会让 loss 曲线更抖,需要更精细的学习率调节。工程上大多数人宁可用 RigL 这种虽粗糙但稳定的方案。

4. 从 SET 到 RigL:动态稀疏算法的演进路线与适用范围

动态稀疏不是一个新概念,它最早的形态可以追溯到 2018 年左右 SET(Sparse Evolutionary Training)的工作。这几年里,算法家族不断壮大,各自适用场景也完全不同。

4.1 SET:随机长回的奠基之作

SET 发布于 2018 年,是最早的完整动态稀疏训练框架。它的规则极其简单:每隔若干轮,按权重绝对值剪掉每层最不重要的部分连接,然后随机生长同样数量的新连接。没错,生长的选择完全是随机的。

这个“随机”看起来像是偷懒,但 SET 的实验结果显示,随机生长已经足以让稀疏网络在 MNIST 和 CIFAR 上接近稠密网络精度。它的意义在于证明了:动态改变稀疏结构这个思路本身是work的,不需要特别花哨的选择策略就能生效。

SET 的局限也很明显:随机生长没有利用梯度信息,在复杂任务和大模型上有瓶颈。而且 SET 没有对“每层保留多少连接”做动态调节,固定配额限制了它在不均衡任务上的上限。今天很少有人在生产环境直接跑 SET,但它的思想被几乎所有后续算法继承。

4.2 RigL:梯度驱动的工业级选择

RigL(Rigging the Lottery)是 2020 年 Google 提出的方法,可以看作动态稀疏领域的“集大成者”。它把剪枝标准从权重幅度升级为“权重×梯度幅度”,同时把随机生长升级为“在梯度最大的未连接位置生长”。

这样做带来的直接收益是:训练轨迹更稳定,收敛速度更快,最终精度远超 SET。RigL 论文里有一个很出名的实验——从零开始训练 90% 稀疏的 ResNet-50,精度几乎和稠密 ResNet-50 持平。这个结果当时让很多人意识到:动态稀疏完全可以替代“先训稠密再剪枝”的常规路线。

工程上我也更推荐 RigL 这套思路:实现不复杂,所有操作都可以在 PyTorch 层面用 mask 操作完成,不需要改底层框架。唯一要注意的是,计算未连接位置的梯度需要一次完整反传,这在高频剪枝时会有额外开销。实际落地的折中是降低剪枝频率,比如每 1000 步剪一次,而不是每 100 步。

4.3 其他值得留意的变体:SNFS、ADP、Dense-Sparse-Dense

  • SNFS(Sparse Networks from Scratch):在 SET 的基础上引入了“梯度累积”作为剪枝指标,也就是把每一轮计算的梯度累加起来代表连接的重要程度。相比单步梯度,累积值更平滑,适合梯度噪声大的任务。
  • ADP(Adaptive Density Pruning):允许每层稀疏度在训练中自适应变化,而不是预设固定配额。做法是把每个层的保留密度也建模成可优化变量,这样资源会自然流向更关键的层。
  • Dense-Sparse-Dense(DSD):有意思的反向思路。它先训练稠密网络,然后剪到稀疏,用稀疏结构训练一段时间,最后再长回稠密网络再训一轮。实验显示这种“压缩-解压”过程能提升最终稠密模型的精度,相当于用稀疏约束做了一次正则化。

这些变体没有绝对优劣,我跟一些同行的经验是:首选 RigL 作为基线,如果任务要求超高稀疏度(>95%),考虑叠加稀疏度升温;如果训练不稳定,再考虑 ADP 这种自适应密度方案。

5. 大模型时代,Dynamic Sparsity 又以新形态回归

聊完经典算法,必须把视角拉回到现在。LLM 时代,动态稀疏不但没有过时,反而有几个方向上重新变得炙手可热。这跟大模型训练和推理的实际瓶颈高度相关——参数多、算力贵、显存有限,稀疏结构是绕不开的优化方向。

5.1 MoE 本质上是稀疏结构的动态路由

混合专家模型(Mixture of Experts,MoE)可能是大家最熟悉的稀疏大模型架构。它把网络分成多个专家子网络,每个 token 只激活其中 Top-K 个专家。这个“每个 token 激活哪些专家”的决策,其实就是一个动态稀疏过程。

传统动态稀疏在权重连接层面做选择,MoE 则把稀疏的单位从“连接”提升到了“子网络”。两者共享同一个核心思想:不是所有参数都需要在每次前向中参与计算,按需激活才是高效之道。

动态稀疏领域里“生长的位置由数据决定”的原则,和 MoE 的“专家选择由 token 决定”在逻辑上一脉相承。所以如果你在传统模型上调过动态稀疏超参,上手 MoE 的时候会发现很多直觉可以直接迁移——比如裁剪掉长期没被路由到的专家,替换成新的随机初始化专家。

5.2 静态稀疏 LLM 的动态补救:KV Cache 与投机采样

当前大模型推理优化里,最头大的其实是 KV Cache 的显存膨胀。上下文越长,KV Cache 占用越大。有些团队开始研究 KV Cache 的“动态稀疏”:不是所有历史 token 对当前 token 的生成都有同等贡献,能不能在推理过程中动态跳过年久失修的 token 的 KV 计算?

这个方向已经有几篇文章在做,核心思路就是根据 attention 分数动态选择参与计算的 KV 子集,把算力聚焦到相关性最高的历史 token 上。注意,这个过程必须在生成过程中实时决策,不能提前固定,因为它与具体解码路径强相关——这天然就是一个动态稀疏问题。

另一个相似的应用是投机采样(Speculative Decoding)里的草稿模型选择。草稿模型和验证模型之间的关系,某种程度上也可以用稀疏化的眼光看:不是每个 token 都需要验证模型全力参与,某些 token 用小模型就能高置信度带过,这就是推理路径上的动态稀疏。

5.3 动态稀疏与大模型预训练结合:训练成本的想象空间

还有一个我在关注的前沿方向:在大模型预训练阶段就引入动态稀疏。目前主流 LLM 预训练都是稠密的,训练完成后才做量化、剪枝、蒸馏等压缩。可如果从第一轮开始就用动态稀疏策略训练,全程只有 60% 到 70% 的参数参与计算,理论上能省下可观的算力和显存。

为什么工业界还没有普遍这么做?最核心的障碍是训练效率:动态稀疏需要周期性计算全局梯度信息并更新 mask,这个操作目前在大规模并行训练框架下并没有高效实现。张量并行、流水线并行的拓扑结构,跟稀疏 mask 的跨设备重排需求天生冲突。很多研究团队正在尝试把 mask 更新过程做成本地化,但距离完全成熟还需要时间。

这个方向一旦跑通,对大模型领域的价值不亚于一次训练框架革命。我们团队目前在做一个小规模验证,初步结果说明 70% 稀疏度的预训练质量尚可,但距离工业级可用还有较大距离。

6. 动态稀疏落地实操记录:踩过的坑与调出的最优配置

最后一部分,分享我在真实业务里落地动态稀疏的完整经验,包括框架选型、超参配置、常见坑点。这些都是踩过之后换来的血泪教训,希望对大家有帮助。

6.1 框架选择:从自己造轮子到依托成熟库

早期我们团队在 PyTorch 上自己写 mask 更新逻辑,核心代码很简单,无非是那几步:计算重要性、sort、剪枝、生长。但真正的复杂度不在算法逻辑,而在和分布式训练的整合。你用 DistributedDataParallel 跑多卡时,mask 要同步到所有卡上,剪枝步骤要确保各卡算出的 mask 一致,否则训练就崩了。

后来我们切换到已有开源库来兜底基础逻辑。推荐两个方向:

  • 逐步淘汰的工具:早期有dynsparsesparsetrain这类研究代码库,常用于复现论文,但维护大多已停滞。
  • 主流的半官方实现:主要依赖torch.nn.utils.prune加上自研的 mask 更新 loop。官方库虽然原生只提供静态剪枝,但它的 mask 管理机制足够干净,动态更新的逻辑可以自行外挂,这个组合目前最稳。

我的建议是:动态稀疏流程不复杂,完全依赖第三方库反而受限制。掌握核心 mask 更新逻辑,然后配合 Pytorch 基础 API 自己维护,是灵活性和维护成本最好的平衡点。

6.2 超参的黄金组合:我总结出的起手配置

以下是我在视觉模型和推荐模型上都验证过的起手超参组合,适合作为第一版跑通基线用:

超参推荐配置依据
剪枝间隔1000 步/次太长则结构僵化,太短则训练不稳
剪枝比例(单次)当前连接的 20%-30%低于 10% 更新太慢,高于 50% 破坏结构
生长策略动量梯度最大稳定性和探索性平衡最好
稀疏度升温线性升温至目标值避免训练早期结构过强约束
学习率调度加入 10% warmup稀疏结构下梯度方差大,需要预热稳定

这里单次剪枝比例特别容易被人忽略。很多人误以为稀疏度 90% 就是每轮剪掉 90%,这完全不对。动态稀疏的“90% 稀疏”是最终状态,单次剪掉比例应该控制在当前存活连接的 20%-30%,然后逐步逼近目标。一次剪太狠,网络来不及适应就废了。

6.3 坑点实录:我的四次典型翻车现场

坑 1:全局剪枝导致 embedding 层被剪光。我们最早在推荐模型上用全局排序做剪枝,跑了两百步后 loss 突然飞升,查了半天发现 embedding 表被剪得只剩 5% 的连接,所有特征都挤在同一维度上。解决方法是 embedding 层固定不剪或单独设置最低保留密度。特别注意,attention 层和 FFN 层的敏感度完全不同,分层管理稀疏度几乎必须。

坑 2:超低稀疏度下 BatchNorm 统计量漂移。高度稀疏的网络,中间层特征分布变化剧烈,BatchNorm 的 running_mean 和 running_var 更新滞后,导致验证集上精度雪崩。解决办法是把 BatchNorm 换成 LayerNorm,或者降低剪枝频率,给 BatchNorm 足够的适应时间。

坑 3:剪枝频率和学习率退火脱节。学习率已经退到很低时,还在高频剪枝,等于不断改变优化目标函数,loss 会出现典型的“锯齿状”不收敛。后来我把剪枝间隔跟学习率调度器联动——学习率每降一个档位,剪枝间隔拉长一倍。训练后期几乎不再改动 mask,让网络专心收敛。

坑 4:和 AMP 混合精度训练的冲突。PyTorch 的自动混合精度会为权重维护 fp32 主副本和 fp16 计算副本。剪枝时只改了 fp16 副本的 mask,但 fp32 主副本里被剪掉的权重值还在,后续优化器更新又把他们“激活”了。这个坑很隐蔽,表现为剪枝后稀疏度显示正确,但几个 epoch 后权重莫名涨回去。解决方案是剪枝时必须同时把 fp32 主副本里对应的权重重置为零,并确保优化器状态也做对应处理。

6.4 性能实测:动态稀疏在推荐模型上的具体收益

最后给出一组我们业务模型的真实数据,方便大家评估这项技术的投入产出比。模型是两层的深度排序网络,原本参数约 1.2 亿,训练数据 5 亿样本。

  • 离线 AUC:稠密基线 0.8021;90% 稀疏动态训练 0.8030;静态剪枝 90% 后微调 0.7976。
  • 推理时延:90% 稀疏模型使用稀疏矩阵乘法库加速后,单请求时延从 12.3ms 降到 7.1ms。
  • 显存占用:训练阶段显存下降约 55%,主要省在优化器状态和梯度存储上。

动态稀疏这套方法论的理解门槛不高,真正难的是把每个组件和你的具体任务对齐。建议上手路径是:先拿 CIFAR 或公开数据集跑通 RigL,熟悉稀疏度、剪枝率、生长策略之间的关系,再放到自己的业务模型上逐步迁移。这类技术对训练稳定性极其敏感,唯有亲手调一遍才能建立直觉。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询