以往的方法是不断地输入数据集,通过反向传播迭代的方法,更新网络权重,从而达到想要的训练结果。这篇论文提供了一个新的角度,对于分类网络来说,首先根据原来的数据集和网络的初始化权重(固定或随机),通过反向传播迭代更新新的数据集(生成的新数据集)。来形成新的几乎等于(大于等于)分类数量的数据集,再经过有限的几次迭代以后就可以达到比较高的精度。

DATASET DISTILLATION 论文总结

新生成的蒸馏数据特别像噪声。

按照论文的章节安排,重点说一下论文的关键部分。

3.1 优化蒸馏数据

文章的主要思路:

(1)原始方法:

DATASET DISTILLATION 论文总结

权重参数随每一次迭代更新。

(2)新方法:

DATASET DISTILLATION 论文总结

利用重新生成的数据集x 和新的学习率η ,通过一次迭代就可以生成想要的参数。更新xη 的公式为:

DATASET DISTILLATION 论文总结

xDATASET DISTILLATION 论文总结ηDATASET DISTILLATION 论文总结 是通过事先梯度迭代学习得到的,是需要初始化的。

这一步讲的是固定参数θ0 ,来更新xDATASET DISTILLATION 论文总结ηDATASET DISTILLATION 论文总结  ,最终得到的准确率效果也比较好,如图1所示,Fixed init一栏就是结果,比其他情况都要好。

DATASET DISTILLATION 论文总结

图1

3.2 权重随机初始化的数据蒸馏

这一小节讲的是权重符合一定分布的情况下的数据蒸馏,这种情况下,精度会有一定的损失,具体方法步骤如图2所示,过程也很清楚。随机初始化权重对应的结果在图1的Random init一栏,效果比Fixed init要差。

DATASET DISTILLATION 论文总结

图2

3.3 一个简单线性案例的分析

本节只是理论推导,任意权重初始化所生成的蒸馏数据集经过一次GD就可以实现最优精度所需要的条件。推倒过程比较简单,这里只列出结论:

DATASET DISTILLATION 论文总结

其中,I是单位矩阵,M是生成的蒸馏数据集的大小,d 是NxD维的数据矩阵,N表示原数据集大小(或者一个minibatch大小),D是分类大小,M<<N。

即要求dTd 是满秩矩阵,并且M≥ D,M至少要等于分类的大小,这也是MNIST只要10个蒸馏数据集图片就可以的原因。当然在随机权重初始化的条件下,这个条件只有在理想情况下才能达到。

3.4 多次GD和多epoch训练

主要是为了实现以下方案——在随机权重初始化条件下,像全数据集那样训练网络:

DATASET DISTILLATION 论文总结

论文中说开发了一种技术叫反向梯度优化(back-gradient optimization),可以极大地加快反向梯度更新,这种技术用到了Hession矩阵,在Pytorch中计算很容易。多epoch,需要输入的蒸馏数据的输入顺序保持不变。与单次迭代效果对比如下图所示:

DATASET DISTILLATION 论文总结

总结

  1. 本论文提供了一个新的视角,先更新数据集再有限地更新权重参数。不过这里的蒸馏跟以往网络压缩范畴的蒸馏技术没有关系。
  2. 计算过程比较繁琐,在新数据集生成过程中,增加了数据集的学习率更新参数,需要更多的计算,并且多epoch情况下还需要蒸馏数据的顺序保持不变。在简单的数据集上比较奏效,数据集复杂(如imagenet等)的情况下,可能计算起来就比较吃力。
  3. 原本认为可以加快网络训练速度,论文也没有提这个方面。

 

相关文章:

  • 2021-06-28
  • 2021-07-21
  • 2021-06-25
  • 2021-09-11
  • 2021-11-27
  • 2021-04-18
  • 2021-10-26
  • 2021-05-29
猜你喜欢
  • 2021-05-09
  • 2021-11-30
  • 2021-07-05
  • 2021-05-12
  • 2021-09-11
  • 2021-11-21
  • 2022-12-23
相关资源
相似解决方案