什么是数据集蒸馏?快速概览
了解数据集蒸馏如何通过用一组小型、经过优化的合成样本替代大型数据集,加快模型训练并降低计算成本。

训练模型似乎是数据科学家工作中最耗时的部分。但他们大部分时间(通常为 60% 到 80%)实际上都花在准备数据上:收集、清理数据,并将其整理为可用于建模的形式。随着数据集规模不断扩大,准备数据所需的时间也会增加,从而拖慢实验进度,让迭代变得更加困难。
为了解决这个问题,研究人员多年来一直在寻找简化训练流程的方法。合成数据、数据集压缩和更优的优化方法等方案,都旨在降低处理大规模数据集的成本和阻力,并加快机器学习工作流。
这引出了一个关键问题:我们能否在大幅缩小数据集的同时,仍然实现与使用完整数据训练模型相同的性能?数据集蒸馏是一个很有前景的答案。
它会创建大型训练数据集的紧凑版本,同时保留模型有效学习所需的关键模式。这样可以加快训练速度、降低计算需求,并提高实验效率。你可以把它想象成模型的学习小抄:一组很小的合成数据示例,旨在教会模型与完整数据集相同的核心模式。
本文将探讨数据集蒸馏的工作原理,以及它如何支持真实应用中的可扩展机器学习和深度学习。让我们开始吧!
理解数据集蒸馏#
数据集蒸馏是这样一个过程:将大型训练数据集压缩成小得多的数据集,同时仍然向模型传授与原始数据集几乎相同的信息。许多研究人员也将这一过程称为数据集凝缩,因为其目标是捕捉完整数据集中呈现的关键模式。
蒸馏数据集不同于随机生成的合成数据,也不同于简单地从真实图像中挑选较小的子集。它不是随机生成的虚假数据集,也不是原始数据集的删减副本。
相反,它经过有针对性的优化,以捕捉最重要的模式。在此过程中,每个像素和特征都会经过调整和优化,使得使用蒸馏数据训练的神经网络几乎能够学到与使用整个数据集训练时相同的内容。
这一想法最早出现在 Tongzhou Wang、Jun-Yan Zhu、Antonio Torralba 和 Alexei A. Efros 于 2018 年发表的一篇 arXiv 论文中。早期测试使用了MNIST和CIFAR-10等简单数据集,因此很容易展示少量蒸馏样本如何替代数千张真实图像。

图 1. 使用数据集蒸馏处理图像数据(来源)
此后,后续研究进一步推动了数据集蒸馏的发展,包括在 ICML 和 ICLR 上发表的方法,使数据凝缩更加高效且可扩展。
数据集蒸馏的意义#
数据集蒸馏可以提高训练效率并加快开发周期。通过减少模型需要学习的数据量,它降低了计算需求。
这对持续学习、神经架构搜索和边缘训练尤其有用:持续学习中的模型会随着时间更新;神经架构搜索会测试大量模型设计;边缘训练则要求模型在内存和电量有限的小型设备上运行。总体而言,这些优势使数据集蒸馏成为许多机器学习工作流中快速初始化、快速微调和构建早期原型的理想选择。
数据集蒸馏的工作原理概览#
数据集蒸馏会创建合成的,或人工生成的训练样本。这些样本帮助模型以接近使用真实数据训练的方式进行学习。其工作过程会在常规训练期间跟踪三个关键因素。
第一是损失函数,即用于表示模型预测错误程度的误差分数。第二是模型参数,即网络在学习过程中不断更新的内部权重。
第三是训练轨迹,它描述误差和权重如何随时间逐步变化。随后,系统会优化合成样本,使模型在其上训练时,误差下降方式和权重更新方式都与使用完整数据集时相同。
逐步了解数据集蒸馏#
下面更详细地了解数据集蒸馏的工作过程:
- **步骤 1 - 初始化合成像素:**流程从充当可学习输入的合成图像开始。起初,这些图像几乎没有结构,看起来像一张白纸。随着时间推移,它们会被优化为包含丰富信息的示例。
- **步骤 2 - 使用梯度匹配和反向传播进行优化:**模型在这些合成图像上训练时,会生成梯度,表示每个像素应如何变化,才能更好地匹配真实数据的训练行为。反向传播是网络从错误中学习的方法。它会将误差反向传过模型,以确定哪些像素和权重导致了误差,然后对其进行微调。利用这些梯度,反向传播会逐步调整合成图像,使其包含更多训练信息。
- **步骤 3 - 匹配训练步骤中的行为:**该方法还会匹配训练轨迹,也就是模型在学习过程中经历的逐步变化。这能确保蒸馏数据集引导模型沿着与使用完整数据集时相似的学习路径前进。
- **步骤 4 - 验证与泛化:**最后,在真实验证数据上评估蒸馏数据集,以了解训练后的模型在新示例上的表现。这可以检查合成数据是否教会模型广泛且有效的模式,而不是导致模型记住特定样本。

图 2. 数据集蒸馏示意(来源)
数据集蒸馏的主要方法#
所有数据集蒸馏方法都建立在同一个核心理念之上,尽管它们可能使用不同的算法来实现。大多数方法可分为三类:性能匹配、分布匹配和参数匹配。
接下来,我们分别了解每种方法及其工作原理。
性能匹配#
数据集蒸馏中的性能匹配,侧重于创建一个经过优化的极小训练集,使模型达到几乎与使用完整原始数据集训练时相同的准确率。蒸馏样本不是随机挑选的,而是经过优化,使得在其上训练的模型能够获得与原始数据集训练模型相似的预测结果、训练过程中的损失表现或最终准确率。
元学习是改进这一过程的常用方法。蒸馏数据集会通过反复的训练回合进行更新,从而在各种可能的情况下都能发挥作用。
在这些回合中,该方法会模拟学生模型如何从当前蒸馏样本中学习,检查学生模型在真实数据上的表现,然后调整蒸馏样本,使其成为更好的教师。随着时间推移,蒸馏数据集会学会支持快速学习和良好泛化,即使学生模型使用不同的初始权重或不同的架构也是如此。这使蒸馏数据集更加可靠,不会局限于单次训练运行。

图 3. 元学习过程(来源)
分布匹配技术#
与此同时,分布匹配会生成与真实数据集统计模式相匹配的合成数据。该方法不只关注模型的最终准确率,还关注神经网络在学习过程中生成的内部特征。
接下来,我们看看推动分布匹配的两种技术。
单层分布匹配#
单层分布匹配专注于神经网络的单个层,并比较该层为真实数据和合成数据生成的特征。这些特征也称为激活值,反映了模型在网络该位置学到的内容。
通过让合成数据生成相似的激活值,该方法促使蒸馏数据集体现与原始数据集相同的重要模式。在实践中,系统会反复更新合成样本,直到所选层的激活值与真实图像产生的激活值高度匹配。
这种方法相对简单,因为它一次只对齐一个层级的表示。对于较小的数据集,或不需要匹配深层多阶段特征层次结构的任务,它尤其有效。通过清晰地对齐一个特征空间,单层匹配为使用蒸馏数据集进行学习提供了稳定且有意义的信号。
多层分布匹配#
多层分布匹配基于比较真实数据和合成数据这一理念,但会在神经网络的多个层执行比较,而不只是一个层。不同层捕捉不同类型的信息:浅层捕捉简单的边缘和纹理,深层则捕捉形状和更复杂的模式。
通过匹配这些层中的特征,蒸馏数据集会被推动去反映模型在多个层级上学到的内容。由于它会对齐整个网络中的特征,该方法可以帮助合成数据保留更丰富的信号,而模型正是依靠这些信号区分不同类别。
这对计算机视觉尤其有帮助,即用于让模型理解图像和视频的任务,因为有用的模式分布在多个层中。当多个深度上的特征分布都能良好匹配时,蒸馏数据集就能更有力、更可靠地替代原始训练数据。
参数匹配方法#
数据集蒸馏中的另一个关键类别是参数匹配。它不匹配准确率或特征分布,而是匹配模型权重在训练过程中的变化方式。通过让模型在蒸馏数据集上的训练产生与真实数据训练相似的参数更新,模型就能沿着几乎相同的学习路径前进。
接下来,我们将介绍两种主要的参数匹配方法。
单步匹配#
单步匹配会比较模型在真实数据上完成一个训练步骤后权重的变化。随后调整蒸馏数据集,使模型在其上训练一步后产生非常相似的权重更新。由于只关注这一次更新,该方法直观且运行速度快。
缺点是,一步训练无法反映完整的学习过程,尤其是在模型需要多次更新才能构建更丰富特征的复杂任务中。因此,单步匹配通常最适合较简单的问题或较小的数据集,因为模型可以快速捕捉有用的模式。
多步参数匹配#
相比之下,多步参数匹配会观察模型权重在多个训练步骤中的变化,而不只是一步。这一系列更新就是模型的训练轨迹。
构建蒸馏数据集时,会使模型在合成样本上训练后的轨迹尽可能接近其在真实数据上训练时的轨迹。通过匹配更长的学习过程,蒸馏数据集能够捕捉原始训练过程中的更多结构。
由于反映了学习随时间展开的方式,多步匹配通常更适合规模更大或更复杂的数据集,因为模型需要多次更新才能捕捉有用模式。它确实需要更多计算,因为必须跟踪多个步骤,但与单步匹配相比,它通常能生成泛化能力更好、性能更高的蒸馏数据集。
合成数据集的生成与优化原理#
在更好地理解主要蒸馏方法后,我们现在可以了解合成数据是如何生成的。在数据集蒸馏中,系统会优化合成样本,使其捕捉最重要的学习信号,从而让小型数据集替代规模大得多的数据集。
接下来,我们将了解这些蒸馏数据如何生成和评估。
创建和评估蒸馏图像#
在数据集蒸馏过程中,合成像素会经过许多训练步骤的更新。神经网络从当前合成图像中学习,并发送基于梯度的反馈,说明每个像素应如何变化,才能更好地匹配真实数据集中的模式。
这一过程之所以有效,是因为它具有可微性(即每个步骤都平滑且具有定义明确的梯度,因此像素的小幅变化会带来可预测的损失变化),从而使模型能够在梯度下降过程中平滑地调整合成数据。
随着优化持续进行,合成图像开始形成有意义的结构,包括模型能够识别的形状和纹理。这些经过优化的合成图像通常用于图像分类任务,因为它们捕捉了分类器需要学习的关键视觉线索。
评估蒸馏数据集时,会使用其训练模型,并将结果与使用真实数据训练的模型进行比较。研究人员会测量验证准确率,并检查合成数据集是否保留了区分类别所需的判别特征(模型依靠其区分不同类别的模式或信号)。他们还会在不同运行或模型设置下测试稳定性和泛化能力,以确保蒸馏数据不会导致过拟合。
数据蒸馏的实际应用#
接下来,我们将详细了解一些示例,看看蒸馏数据集如何在数据有限或高度专业化的情况下加快训练、降低计算成本,同时保持较强的性能。
使用数据集蒸馏处理计算机视觉应用#
在计算机视觉领域,目标是训练模型理解图像和视频等视觉数据。这些模型会学习边缘、纹理、形状和物体等模式,然后将其用于图像分类、目标检测或分割等任务。由于视觉问题通常在光照、背景和视角方面存在巨大变化,计算机视觉系统通常需要大型数据集才能实现良好泛化,因此训练成本高且速度慢。

图 4. 数据集蒸馏示例(来源)
在医学扫描、野生动物监测或工厂缺陷检测等图像分类应用中,模型通常面临准确率与训练成本之间的艰难权衡。这些任务通常涉及海量数据集。
数据集蒸馏可以将原始训练集压缩为少量合成图像,同时保留分类器所需的最重要视觉线索。在 ImageNet 等大型基准测试中,研究表明,仅使用约原始图像的 4.2%的蒸馏数据集也能保持较高的分类准确率。这意味着一个极小的合成代理数据集就能以低得多的计算成本替代数百万个真实样本。
神经架构搜索#
神经架构搜索(NAS)是一种自动探索多种神经网络设计,以找到最适合某项任务的设计的技术。由于 NAS 必须训练和评估大量候选模型,在完整数据集上运行可能非常缓慢且计算密集。
数据集蒸馏通过创建一个仍保留原始数据主要学习信号的极小合成训练集来提供帮助,因此可以更快地测试每个候选架构。这让 NAS 能够高效比较不同设计,同时相对可靠地保持优劣架构的排序,在不大幅牺牲最终模型质量的情况下降低搜索成本。
持续学习与边缘部署#
持续学习系统,也就是随着新数据到来而持续更新、而不是只训练一次的模型,需要快速且节省内存的更新。边缘设备(如摄像头、手机和传感器)也面临类似限制,因为其计算和存储预算都很紧张。
数据集蒸馏通过将大型训练集压缩成极小的合成数据集,在这两种场景中都能发挥作用,使模型可以使用小型回放集而不是完整数据集进行适应或重新训练。例如,基于核的元学习研究表明,仅使用 10 个蒸馏样本,就能在 CIFAR-10 这一标准图像分类基准上达到超过 64% 的准确率。由于回放集非常紧凑,更新速度会快得多,也更加实用,尤其适合需要频繁刷新模型的场景。
数据集蒸馏还可以与大型语言模型的知识蒸馏协同工作。小型蒸馏数据集可以保留教师模型最重要的任务信号,使压缩后的学生模型能够更高效地训练或更新,同时不会损失太多性能。由于这些数据集非常小,它们特别适合边缘端或设备端使用:存储和计算资源有限,但你仍希望模型在更新后保持准确。
数据蒸馏的优缺点#
以下是使用数据集蒸馏的一些优势:
- **非常适合快速实验。**你可以测试新的架构、损失函数或超参数,而不必每次都在海量数据集上重新训练。
- **可能带来隐私优势。**共享蒸馏合成样本可能比共享真实用户数据点更安全,因为原始示例不会被直接暴露。
- **通常优于简单挑选子集。**蒸馏不是单纯选择示例,而是主动优化示例,使其包含尽可能多的信息。
虽然数据集蒸馏具有多项优势,但也有一些需要注意的局限:
- 过拟合**:**蒸馏数据通常最适合蒸馏过程中使用的架构,迁移到差异很大的模型时可能表现不佳。
- **对超参数敏感。**结果可能在很大程度上取决于学习率、初始化方式或蒸馏步骤数等因素。
- **更难扩展到真实世界的复杂性。**在基准测试上表现良好的方法,在规模大、数据混杂或分辨率高的数据集上可能会损失准确率。
要点总结#
数据集蒸馏使少量合成样本能够几乎像完整数据集一样有效地教会模型。这让机器学习变得更快、更高效,也更易于扩展。随着模型不断增大并需要更多数据,蒸馏数据集提供了一种切实可行的方法,可以在不牺牲准确率的情况下降低计算成本。
加入我们的社区,并查看我们的 GitHub 仓库,了解更多 AI 相关内容。如果你想构建自己的视觉 AI 项目,可以查看我们的许可选项。访问我们的解决方案页面,进一步了解医疗保健领域的 AI和零售领域的视觉 AI等应用。









