Skip to main content
QUICK REVIEW

[论文解读] Nonparametric Bayesian Deep Networks with Local Competition

Konstantinos P. Panousis, Sotirios Chatzis|arXiv (Cornell University)|May 19, 2018
Gaussian Processes and Bayesian Inference参考文献 29被引用 3
一句话总结

该论文提出SB-LWTA网络,这是一种非参数贝叶斯深度学习框架,通过使用局部赢家通吃(LWTA)非线性激活和棒棒糖先验,实现在训练过程中对最小网络复杂度和最优浮点精度的推断。通过利用离散隐变量建模组件效用并执行贝叶斯推断,该方法在保持高精度的同时,显著降低了计算开销和预测时间,在MNIST和LeNet-5-Caffe基准测试中优于现有方法。

ABSTRACT

The aim of this work is to enable inference of deep networks that retain high accuracy for the least possible model complexity, with the latter deduced from the data during inference. To this end, we revisit deep networks that comprise competing linear units, as opposed to nonlinear units that do not entail any form of (local) competition. In this context, our main technical innovation consists in an inferential setup that leverages solid arguments from Bayesian nonparametrics. We infer both the needed set of connections or locally competing sets of units, as well as the required floating-point precision for storing the network parameters. Specifically, we introduce auxiliary discrete latent variables representing which initial network components are actually needed for modeling the data at hand, and perform Bayesian inference over them by imposing appropriate stick-breaking priors. As we experimentally show using benchmark datasets, our approach yields networks with less computational footprint than the state-of-the-art, and with no compromises in predictive accuracy.

研究动机与目标

  • 通过实现数据驱动的模型复杂度推断,解决深度神经网络中的过参数化和高计算成本问题。
  • 通过自动网络剪枝和精度压缩,减少模型冗余,提升在资源受限设备上的可扩展性。
  • 构建一个基于非参数先验的系统性贝叶斯框架,联合推断网络结构与参数精度。
  • 利用局部赢家通吃(LWTA)机制,实现生物上合理、稀疏且具有判别性的表征。
  • 在预测精度和计算效率两方面,超越现有的正则化、蒸馏和剪枝方法。

提出的方法

  • 提出一种基于局部赢家通吃(LWTA)单元的深度网络架构,通过侧向抑制机制确保每组内仅一个单元处于激活状态。
  • 引入辅助离散隐变量,表示哪些网络组件(单元或连接)实际用于数据建模。
  • 在这些隐变量上应用棒棒糖先验,以实现对组件效用和模型复杂度的非参数贝叶斯推断。
  • 使用随机梯度变分贝叶斯(SGVB)方法,实现对网络组件和权重精度的高效训练与后验推断。
  • 通过分析推断出的权重后验方差,推断最优浮点精度,实现在不损失精度的前提下实现压缩。
  • 将SB-LWTA模型构建为启发式蒸馏与正则化技术的系统性、数据驱动替代方案。

实验结果

研究问题

  • RQ1贝叶斯非参数框架能否有效推断给定数据集所需的最小网络复杂度?
  • RQ2在深度网络中使用局部竞争(LWTA)是否能带来比标准非线性激活更高效、更准确的模型?
  • RQ3棒棒糖先验能否被有效适配以同时推断深度网络的结构与参数精度?
  • RQ4所提出方法在多大程度上降低了计算开销,同时保持或提升预测精度?
  • RQ5LWTA模块中的胜出选择模式在不同类别间如何反映判别性、可泛化的特征?

主要发现

  • 在MNIST数据集上,SB-LWTA网络在所有对比方法中取得了最高的预测精度,甚至优于原始的LeNet-5-Caffe架构。
  • 该方法将特征图数量减少至所有基线中的最低水平,表明其具有优越的结构压缩能力。
  • 与原始网络相比,预测时间降低了整整一个数量级,证明了显著的推理效率提升。
  • 每轮训练的平均时间仅比原始网络高出10%,表明训练开销极低。
  • 不同MNIST数字的胜出单元概率存在显著差异,表明模型学习到了具有类别特异性的判别模式。
  • 不同数字对之间胜出单元的重叠率始终低于50%,证实赢家通吃机制编码了独特且可泛化的表征。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。