以往的方法是不断地输入数据集,通过反向传播迭代的方法,更新网络权重,从而达到想要的训练结果。这篇论文提供了一个新的角度,对于分类网络来说,首先根据原来的数据集和网络的初始化权重(固定或随机),通过反向传播迭代更新新的数据集(生成的新数据集)。来形成新的几乎等于(大于等于)分类数量的数据集,再经过有限的几次迭代以后就可以达到比较高的精度。
新生成的蒸馏数据特别像噪声。
按照论文的章节安排,重点说一下论文的关键部分。
3.1 优化蒸馏数据
文章的主要思路:
(1)原始方法:
权重参数随每一次迭代更新。
(2)新方法:
利用重新生成的数据集x 和新的学习率η ,通过一次迭代就可以生成想要的参数。更新x 和η 的公式为:
x 和η
是通过事先梯度迭代学习得到的,是需要初始化的。
这一步讲的是固定参数θ0 ,来更新x 和η
,最终得到的准确率效果也比较好,如图1所示,Fixed init一栏就是结果,比其他情况都要好。
图1
3.2 权重随机初始化的数据蒸馏
这一小节讲的是权重符合一定分布的情况下的数据蒸馏,这种情况下,精度会有一定的损失,具体方法步骤如图2所示,过程也很清楚。随机初始化权重对应的结果在图1的Random init一栏,效果比Fixed init要差。
图2
3.3 一个简单线性案例的分析
本节只是理论推导,任意权重初始化所生成的蒸馏数据集经过一次GD就可以实现最优精度所需要的条件。推倒过程比较简单,这里只列出结论:
其中,I是单位矩阵,M是生成的蒸馏数据集的大小,d 是NxD维的数据矩阵,N表示原数据集大小(或者一个minibatch大小),D是分类大小,M<<N。
即要求dTd 是满秩矩阵,并且M≥ D,M至少要等于分类的大小,这也是MNIST只要10个蒸馏数据集图片就可以的原因。当然在随机权重初始化的条件下,这个条件只有在理想情况下才能达到。
3.4 多次GD和多epoch训练
主要是为了实现以下方案——在随机权重初始化条件下,像全数据集那样训练网络:
论文中说开发了一种技术叫反向梯度优化(back-gradient optimization),可以极大地加快反向梯度更新,这种技术用到了Hession矩阵,在Pytorch中计算很容易。多epoch,需要输入的蒸馏数据的输入顺序保持不变。与单次迭代效果对比如下图所示:
总结
- 本论文提供了一个新的视角,先更新数据集再有限地更新权重参数。不过这里的蒸馏跟以往网络压缩范畴的蒸馏技术没有关系。
- 计算过程比较繁琐,在新数据集生成过程中,增加了数据集的学习率更新参数,需要更多的计算,并且多epoch情况下还需要蒸馏数据的顺序保持不变。在简单的数据集上比较奏效,数据集复杂(如imagenet等)的情况下,可能计算起来就比较吃力。
- 原本认为可以加快网络训练速度,论文也没有提这个方面。