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

发布时间:2026/7/24 15:48:24
基于彩票假说的模型剪枝:找到那个中奖的子网络 基于彩票假说的模型剪枝找到那个中奖的子网络彩票假说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$但仅在$m1$的位置保留参数训练得到的子网络$f(x; m \odot \theta_T)$的性能不低于原网络。关键约束是子网络必须从相同的随机初始化开始训练。如果使用不同的随机种子初始化子网络这个中奖特性就会消失——这暗示着初始化和结构之间的耦合关系是彩票假说的核心。二、迭代式幅度剪枝的完整实现IMPIterative Magnitude Pruning是验证彩票假说的标准算法。其核心操作为随机初始化网络保存初始参数$\theta_0$完整训练网络T步得到$\theta_T$按$|\theta_T|$降序排列保留前p%的参数其余参数对应的掩码设为0将存活参数的权重重置为$\theta_0$中的对应值在新的掩码约束下重新训练重复步骤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, dtypetorch.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_tupleTrue) 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, largestFalse ) # 将对应的掩码位置设为 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): 将存活权重重置为初始随机值彩票假说的关键步骤。 注意仅重置 mTrue 的权重被剪掉的权重保持为零 不参与后续训练。 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上进行了验证实验使用MRPCMicrosoft Research Paraphrase Corpus作为下游任务。实验设置了三组对照IMP 初始化重置标准彩票假说流程每次剪枝后保留原始初始化IMP 随机初始化每次剪枝后使用新的随机种子重新初始化随机剪枝 初始化重置随机选择要剪掉的权重而非按幅值结果达到90%稀疏度时方法MRPC F1相比稠密模型的Δ稠密BERT-base基线88.9—IMP 初始化重置LTH87.2-1.7IMP 随机初始化84.1-4.8随机剪枝 初始化重置82.5-6.4IMP初始化重置组合的精度损失显著小于其他组合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和权重反刍等修正表明早期训练的稳定化对发现子网络同样重要。从工程角度看彩票假说提供了一种超越压缩视角的剪枝方法论——剪枝不仅是去除冗余更是发现精华。