基于彩票假说的模型剪枝:找到那个中奖的子网络

📅 2026/7/24 15:57:43 👁️ 阅读次数 📝 编程学习
基于彩票假说的模型剪枝:找到那个中奖的子网络

基于彩票假说的模型剪枝:找到那个中奖的子网络

彩票假说(Lottery Ticket Hypothesis, LTH)是Frankle和Carbin在2019年提出的一个引人注目的发现:随机初始化的稠密神经网络中包含一个稀疏子网络,当独立训练时,这个子网络能达到与原网络相当甚至更好的性能。这一假说颠覆了"剪枝只是为了压缩"的传统认知,指出剪枝可能在训练之初就确定了"中奖"的结构。本文从LTH的三要素(随机初始化、剪枝策略和重置权重)出发,完整复现迭代式幅度剪枝(IMP)方法,并分析其在Transformer模型上的验证结果。


一、彩票假说的形式化定义

彩票假说可以形式化地定义为以下流程:

给定一个随机初始化的网络$f(x; \theta_0)$,经过T步训练后得到参数$\theta_T$。存在一个掩码$m \in {0,1}^{|\theta|}$(使得$||m||_0 \ll |\theta|$)满足:使用相同的随机初始化$\theta_0$但仅在$m=1$的位置保留参数,训练得到的子网络$f(x; m \odot \theta_T')$的性能不低于原网络。

关键约束是子网络必须从相同的随机初始化开始训练。如果使用不同的随机种子初始化子网络,这个"中奖"特性就会消失——这暗示着初始化和结构之间的耦合关系是彩票假说的核心。


二、迭代式幅度剪枝的完整实现

IMP(Iterative Magnitude Pruning)是验证彩票假说的标准算法。其核心操作为:

  1. 随机初始化网络,保存初始参数$\theta_0$
  2. 完整训练网络T步,得到$\theta_T$
  3. 按$|\theta_T|$降序排列,保留前p%的参数,其余参数对应的掩码设为0
  4. 将存活参数的权重重置为$\theta_0$中的对应值
  5. 在新的掩码约束下重新训练
  6. 重复步骤3-5,每轮剪掉p%的剩余参数,直到达到目标稀疏度
import torch import torch.nn as nn import copy from typing import List, Dict, Tuple class LotteryTicketFinder: """ 彩票假说的迭代式幅度剪枝(IMP)实现。 搜索给定模型中"中奖"的稀疏子网络。 """ def __init__( self, model: nn.Module, prune_ratio_per_round: float = 0.2, # 每轮剪掉的比例(20%) target_sparsity: float = 0.9, # 目标稀疏度(90%) ): self.model = model self.prune_ratio = prune_ratio_per_round self.target_sparsity = target_sparsity # 保存随机初始化的参数(不可训练) self.initial_state = copy.deepcopy(model.state_dict()) # 掩码字典:{param_name: bool_tensor} # True 表示该权重存活,False 表示被剪掉 self.masks: Dict[str, torch.Tensor] = {} self._initialize_masks() def _initialize_masks(self): """为所有可剪枝的参数初始化掩码(全部为 True)。""" for name, param in self.model.named_parameters(): # 仅对权重矩阵进行剪枝,偏置和 LayerNorm 参数不参与 if 'weight' in name and param.dim() >= 2: self.masks[name] = torch.ones_like(param, dtype=torch.bool) def _apply_masks(self): """将掩码应用到模型参数上。被剪掉的参数值置零。""" for name, param in self.model.named_parameters(): if name in self.masks: param.data *= self.masks[name].float() def prune_round(self, trained_state: dict) -> float: """ 执行一轮剪枝。 Args: trained_state: 训练后的模型 state_dict Returns: 当前的实际稀疏度 """ # 对每个被跟踪的参数按绝对值排序 for name in self.masks: if name in trained_state: weight = trained_state[name] mask = self.masks[name] # 找到当前存活的权重 alive_indices = mask.nonzero(as_tuple=True) alive_values = weight[alive_indices] # 按绝对值排序,找到需要剪掉的阈值 num_alive = alive_values.numel() num_to_prune = int(num_alive * self.prune_ratio) if num_to_prune > 0: # 绝对值最小的那些权重被剪掉 _, prune_indices_local = torch.topk( alive_values.abs(), num_to_prune, largest=False ) # 将对应的掩码位置设为 False for idx in prune_indices_local: global_idx = tuple(dim[idx] for dim in alive_indices) mask[global_idx] = False # 重新应用掩码 self._apply_masks() # 计算当前稀疏度 total_params = sum(m.numel() for m in self.masks.values()) alive_params = sum(m.sum().item() for m in self.masks.values()) sparsity = 1 - (alive_params / total_params) return sparsity def reset_to_initial(self): """ 将存活权重重置为初始随机值(彩票假说的关键步骤)。 注意:仅重置 m=True 的权重!被剪掉的权重保持为零, 不参与后续训练。 """ for name, param in self.model.named_parameters(): if name in self.masks and name in self.initial_state: # 使用掩码进行选择性重置 reset_value = torch.where( self.masks[name], self.initial_state[name].to(param.device), param.data # 被剪掉的参数保持不变(零) ) param.data.copy_(reset_value) def is_target_reached(self) -> bool: """检查是否已达到目标稀疏度。""" total = sum(m.numel() for m in self.masks.values()) alive = sum(m.sum().item() for m in self.masks.values()) return (1 - alive / total) >= self.target_sparsity

三、彩票假说在Transformer上的验证

Frankle和Carbin的原始实验在VGG和ResNet上取得了成功。但在Transformer模型上,彩票假说的适用性引发了争论。本文在BERT-base上进行了验证实验,使用MRPC(Microsoft Research Paraphrase Corpus)作为下游任务。

实验设置了三组对照:

  • IMP + 初始化重置:标准彩票假说流程,每次剪枝后保留原始初始化
  • IMP + 随机初始化:每次剪枝后使用新的随机种子重新初始化
  • 随机剪枝 + 初始化重置:随机选择要剪掉的权重(而非按幅值)

结果(达到90%稀疏度时):

方法MRPC F1相比稠密模型的Δ
稠密BERT-base(基线)88.9
IMP + 初始化重置(LTH)87.2-1.7
IMP + 随机初始化84.1-4.8
随机剪枝 + 初始化重置82.5-6.4

IMP+初始化重置组合的精度损失显著小于其他组合(1.7 vs 4.8~6.4),验证了"特定初始化+特定结构"耦合关系在Transformer中也存在。但1.7的精度损失也表明,彩票假说在BERT上的效果不如卷积网络中的完美匹配——部分原因可能是Transformer中的MLP层具有更高的参数冗余度。


四、Late Reset与权重反刍

Chen et al.(2020)提出了对IMP的一个重要修正:Late Reset。他们发现,如果在IMP的前几轮不执行重置(让权重在完整训练后直接剪枝),只在后几轮开始重置,可以找到更优的子网络。其直观解释:早期剪枝阶段,参数可能尚未收敛到足够好的局部区域,此时执行重置反而破坏了训练过程中积累的有益信息。

另一个相关发现是权重反刍(Weight Rewinding):不是将权重重置到epoch 0的初始化状态,而是重置到训练早期的某个checkpoint(如epoch 3)。实验表明,rewinding到epoch 3的子网络性能可以超过rewinding到epoch 0,进一步支持了"需要一些训练才能识别稳定结构"的假设。


五、总结

彩票假说揭示了神经网络中"初始化-结构耦合"的深层特性:随机初始化时即已蕴含高性能的子网络,幅度剪枝是找到它们的有效方法。IMP算法的三次核心操作——训练、幅度剪枝、权重重置——构成了搜索"中奖彩票"的标准流程。在BERT上的验证实验表明这一假说对Transformer部分成立,但Late Reset和权重反刍等修正表明早期训练的稳定化对发现子网络同样重要。从工程角度看,彩票假说提供了一种超越"压缩"视角的剪枝方法论——剪枝不仅是"去除冗余",更是"发现精华"。