ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

动态稀疏:从剪枝到训练范式重塑神经网络高效训练

2026/9/23 6:47:51 拓冰建站 浏览量
动态稀疏:从剪枝到训练范式重塑神经网络高效训练 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 动态稀疏的循环结构剪枝-生长-再训练动态稀疏的训练流程是一个不断循环的结构核心有三个阶段循环往复直到训练结束适度训练在当前的稀疏 mask 下训练若干步让现有连接对应的权重充分收敛。剪枝按照重要性指标移除一部分当前“最不重要”的连接释放出对应的稀疏度预算。生长在释放的预算里按某种策略选择新的连接位置把 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 变化幅度小收敛稳定。如果power1线性增长训练后期 mask 变化还是太大loss 容易在最后阶段翘尾回升。3.4 DSR 的重参数化技巧让稀疏 mask 直接参与反向传播上面讲的剪枝和生长都是在离散 mask 上操作这个过程不可导。大多数动态稀疏算法都是把 mask 当“开关”用梯度不经过 mask 本身这也意味着网络没法通过学习来调整“哪些连接重要”这种高层行为。Deep Sparse RewiringDSR提出的重参数化思路解决的就是这个问题。它的做法是不给每个权重二进制的 0/1 mask而是给每个权重一个连续分布参数这个参数控制该连接“存在”的概率。训练时从分布中采样出实际 mask采样过程用 Gumbel-Softmax 之类的连续近似替代这样 mask 的“概率参数”就能参与反向传播。换句话说普通动态稀疏是“外力”决定谁死谁活DSR 是让网络自己学习该让谁死谁活。DSR 的核心数学形式类似变分推断——每个连接的重要性被建模成一个可学习的分布每次前向采样得到的稀疏结构天然带了随机性相当于在做结构层面的数据增强。不过说实话DSR 在中小规模模型上效果不错但在大规模训练里工程复杂度偏高。一个原因是它对每个权重都要额外维护分布参数显存开销翻倍另一个原因是采样带来的随机性会让 loss 曲线更抖需要更精细的学习率调节。工程上大多数人宁可用 RigL 这种虽粗糙但稳定的方案。4. 从 SET 到 RigL动态稀疏算法的演进路线与适用范围动态稀疏不是一个新概念它最早的形态可以追溯到 2018 年左右 SETSparse Evolutionary Training的工作。这几年里算法家族不断壮大各自适用场景也完全不同。4.1 SET随机长回的奠基之作SET 发布于 2018 年是最早的完整动态稀疏训练框架。它的规则极其简单每隔若干轮按权重绝对值剪掉每层最不重要的部分连接然后随机生长同样数量的新连接。没错生长的选择完全是随机的。这个“随机”看起来像是偷懒但 SET 的实验结果显示随机生长已经足以让稀疏网络在 MNIST 和 CIFAR 上接近稠密网络精度。它的意义在于证明了动态改变稀疏结构这个思路本身是work的不需要特别花哨的选择策略就能生效。SET 的局限也很明显随机生长没有利用梯度信息在复杂任务和大模型上有瓶颈。而且 SET 没有对“每层保留多少连接”做动态调节固定配额限制了它在不均衡任务上的上限。今天很少有人在生产环境直接跑 SET但它的思想被几乎所有后续算法继承。4.2 RigL梯度驱动的工业级选择RigLRigging the Lottery是 2020 年 Google 提出的方法可以看作动态稀疏领域的“集大成者”。它把剪枝标准从权重幅度升级为“权重×梯度幅度”同时把随机生长升级为“在梯度最大的未连接位置生长”。这样做带来的直接收益是训练轨迹更稳定收敛速度更快最终精度远超 SET。RigL 论文里有一个很出名的实验——从零开始训练 90% 稀疏的 ResNet-50精度几乎和稠密 ResNet-50 持平。这个结果当时让很多人意识到动态稀疏完全可以替代“先训稠密再剪枝”的常规路线。工程上我也更推荐 RigL 这套思路实现不复杂所有操作都可以在 PyTorch 层面用 mask 操作完成不需要改底层框架。唯一要注意的是计算未连接位置的梯度需要一次完整反传这在高频剪枝时会有额外开销。实际落地的折中是降低剪枝频率比如每 1000 步剪一次而不是每 100 步。4.3 其他值得留意的变体SNFS、ADP、Dense-Sparse-DenseSNFSSparse Networks from Scratch在 SET 的基础上引入了“梯度累积”作为剪枝指标也就是把每一轮计算的梯度累加起来代表连接的重要程度。相比单步梯度累积值更平滑适合梯度噪声大的任务。ADPAdaptive Density Pruning允许每层稀疏度在训练中自适应变化而不是预设固定配额。做法是把每个层的保留密度也建模成可优化变量这样资源会自然流向更关键的层。Dense-Sparse-DenseDSD有意思的反向思路。它先训练稠密网络然后剪到稀疏用稀疏结构训练一段时间最后再长回稠密网络再训一轮。实验显示这种“压缩-解压”过程能提升最终稠密模型的精度相当于用稀疏约束做了一次正则化。这些变体没有绝对优劣我跟一些同行的经验是首选 RigL 作为基线如果任务要求超高稀疏度95%考虑叠加稀疏度升温如果训练不稳定再考虑 ADP 这种自适应密度方案。5. 大模型时代Dynamic Sparsity 又以新形态回归聊完经典算法必须把视角拉回到现在。LLM 时代动态稀疏不但没有过时反而有几个方向上重新变得炙手可热。这跟大模型训练和推理的实际瓶颈高度相关——参数多、算力贵、显存有限稀疏结构是绕不开的优化方向。5.1 MoE 本质上是稀疏结构的动态路由混合专家模型Mixture of ExpertsMoE可能是大家最熟悉的稀疏大模型架构。它把网络分成多个专家子网络每个 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 一致否则训练就崩了。后来我们切换到已有开源库来兜底基础逻辑。推荐两个方向逐步淘汰的工具早期有dynsparse、sparsetrain这类研究代码库常用于复现论文但维护大多已停滞。主流的半官方实现主要依赖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.802190% 稀疏动态训练 0.8030静态剪枝 90% 后微调 0.7976。推理时延90% 稀疏模型使用稀疏矩阵乘法库加速后单请求时延从 12.3ms 降到 7.1ms。显存占用训练阶段显存下降约 55%主要省在优化器状态和梯度存储上。动态稀疏这套方法论的理解门槛不高真正难的是把每个组件和你的具体任务对齐。建议上手路径是先拿 CIFAR 或公开数据集跑通 RigL熟悉稀疏度、剪枝率、生长策略之间的关系再放到自己的业务模型上逐步迁移。这类技术对训练稳定性极其敏感唯有亲手调一遍才能建立直觉。