利用多臂老虎机进行卷积神经网络中的损失感知特征图剪枝
本文提出一种基于多臂老虎机的损失感知特征图剪枝框架,通过将每个特征图视为一个臂,临时屏蔽并评估损失变化,从而在保持精度的前提下减少冗余特征图和计算量。实验表明,UCB1和Thompson Sampling方法在多个数据集上性能优于贪心剪枝和基于幅度的剪枝,且与未剪枝模型统计上相当。
卷积神经网络(CNN)通常包含大量冗余特征图,这些冗余特征图不仅增加了模型的存储开销,还显著提升了推理计算成本。为了解决这一问题,本文提出了一种基于多臂老虎机(Multi-Armed Bandits)的损失感知特征图剪枝框架。该框架的核心思想是将每个候选特征图视为一个“臂”,通过动态评估移除每个特征图对损失函数的影响,从而做出更明智的剪枝决策。具体而言,在每个操作回合中,算法会临时屏蔽一个特征图,并在一个小批量样本上计算损失变化。之后恢复该特征图,并将观察到的损失变化转换为一个“安全移除奖励”。在固定的操作预算(play budget)完成后,所有候选特征图会根据学习到的分数进行排序,得分最高的k个特征图将被永久移除,同时移除对应的滤波器、偏置以及下一层的输入通道核。这种结构化剪枝方式直接减少了卷积层的输出通道数,从而在保持模型精度的前提下显著降低计算量。
该研究评估了两种经典的多臂老虎机策略:UCB1(Upper Confidence Bound)和Thompson Sampling。实验首先在LeNet网络和MNIST数据集上进行直接/或acle式比较,随后扩展到更复杂的任务和数据集,包括MNIST、CIFAR-10、CIFAR-100、SVHN、CUB-200-2011以及Oxford Flowers 102。结果表明,UCB1和Thompson Sampling在移除大量特征图并减少卷积计算量的同时,能够保持与未剪枝模型非常接近的准确率。通过Friedman秩和检验和Nemenyi事后检验,研究进一步证实:UCB1在所有候选方法中取得了最高的平均秩,Thompson Sampling紧随其后;两者均显著优于传统的贪心剪枝(greedy pruning)和基于幅度的剪枝(magnitude-based pruning),同时在统计上与原始未剪枝模型的表现相当。这表明,基于多臂老虎机的剪枝方法在自动化程度和最终模型性能之间取得了良好的平衡。
该工作的创新之处在于将损失感知引入特征图剪枝,利用强化学习中的探索-利用权衡来指导剪枝过程,避免了传统方法中因忽视特征图重要性差异而导致的精度下降。此外,该方法无需预训练或微调阶段,直接在原始模型上进行评估,大大简化了剪枝流程。这项研究为深度学习模型的压缩提供了新的视角。在资源受限的部署场景中,如移动设备和嵌入式系统,卷积神经网络的高效推理至关重要。本文提出的损失感知剪枝框架不仅减少了存储占用和计算延迟,还为自动化模型优化开辟了新的路径。未来,该方法可能被应用于更大规模的模型和更复杂的网络架构,为神经网络的压缩与加速提供新的思路。