通过影响匹配进行数据集蒸馏 Dataset Distillation by Influence Matching
提出可微影响估计器,让合成数据对最终模型参数的影响直接匹配全量数据
前置知识
数据集蒸馏
给定大规模真实数据集 $D$,学习一个极小的合成集 $S$(每类仅 1–50 张图,即 IPC 设置),使得在 $S$ 上训练的模型在真实测试集上的精度尽可能接近在 $D$ 上训练的模型。其原始形式是双层优化:内层在 $S$ 上训练网络得到 $\theta^*_S$,外层优化 $S$ 使 $\theta^*_S$ 在 $D$ 上损失最小,直接求解非常困难。
这是本文要解决的任务本身,理解其双层优化结构与 IPC 评估协议,才能看懂 Inf-Match 的目标函数设计和实验设置。
影响函数
量化「训练集中加入或移除某些样本会对训好后的模型参数产生多大改变」的工具。经典做法(Koh & Liang, 2017)用逆 Hessian 乘梯度来近似留一重训练的影响,但要求损失对参数凸,且逆 Hessian 计算在大模型上代价高昂,随后有 TracIn、DataInf 等改进但仍受限制。
本文方法名为「影响匹配」,其核心贡献正是绕开凸性假设与逆 Hessian 计算,构造出线性时间、完全可微的样本级影响估计器,不懂影响函数就无法理解其目标函数。
梯度匹配与轨迹匹配
两类主流「过程匹配」蒸馏方法:梯度匹配(GM/DSA)对齐真实与合成数据在训练每一步产生的梯度;轨迹匹配(MTT 及其改进 DATM)让合成数据引导的参数状态与真实训练轨迹一致。它们都用中间优化信号作为最终性能的启发式代理。
它们是本文的对照范式与主要实验基线(NCFM、DATM 等),论文的核心动机正是「过程一致并不蕴含结果一致」这一优化鸿沟。
Hessian-向量积与 Pearlmutter 技巧
不显式构造 Hessian 矩阵而计算 $H \cdot G$ 的方法:利用 $\nabla_\theta L(\theta+\epsilon G) \approx \nabla_\theta L(\theta) + \epsilon H G$,一次额外反向传播即可以 $\mathcal{O}(p)$ 复杂度得到结果(论文式 4),$p$ 为参数量,主流深度学习框架都直接支持。
Inf-Match 的影响估计器中充满 $H^t_D G^t_Z$ 这类 Hessian-梯度乘积项,正是靠该技巧才做到线性时间且完全免逆 Hessian。
软标签蒸馏
用训练好模型的输出概率分布(soft label)而非独热码作为监督信号,能够携带类间相似性信息、提高表示效率。数据集蒸馏中常把合成图像的标签也设为可学习变量,与图像联合优化。
本文初始化合成集时用最终模型 $\theta^T_D$ 生成软标签并联合更新,消融实验显示可学习标签是达到 57.4% 最终精度的关键组件之一,是复现方法的重要细节。
研究动机
视觉数据规模已膨胀到数千万乃至数亿样本,存储、传输和训练成本高企,数据集蒸馏因此成为关键任务:用极小合成集尽可能保留全量数据的训练效果。然而现有主流方法都以「过程代理」为目标:梯度匹配(GM)对齐真实与合成数据每一步的梯度,轨迹匹配(MTT/DATM)对齐训练轨迹,NCFM 从极小极大视角匹配神经特征函数。这类启发式假设存在根本缺陷——合成数据可以完美复现中间训练行为,却不能保证最终精度与泛化,形成「优化鸿沟」。若想直接对齐最终结果,就需要量化单个样本或子集对最终模型参数的影响,但经典影响函数方法(如 Koh & Liang 2017 及其后续 DataInf 等)依赖损失函数的凸性假设,且需要计算耗时的逆 Hessian-梯度乘积,在深度非凸网络上既不成立也不可行。因此学界长期被迫在「计算可行性」(过程对齐)与「目标忠实度」(结果对齐)之间取舍。
本文的目标是本文的目标是把数据集蒸馏的优化目标从「模仿训练过程」改为「对齐训练结果」:学习一个小型合成集 $S$,使得在其上训练得到的最终模型参数与在全量数据 $D$ 上训练的结果尽可能一致。具体而言,作者定义了移除影响 $I^-_Z = \theta^*_{D-Z} - \theta^*_D$ 与添加影响 $I^+_Z = \theta^*_{D+Z} - \theta^*_D$,并提出一个完全可微、线性时间、无凸性假设、无需逆 Hessian 计算的样本级影响估计器(定理 1),随后通过最小化 $\|I^-_D + I^+_S\|$ 来学习合成集,使 $S$ 的添加影响恰好抵消移除 $D$ 的影响。最终在 CIFAR-10、CIFAR-100、Tiny-ImageNet 分类基准上全面超越 DATM、NCFM 等强基线,并成功扩展到 Flickr30K 上的视觉-语言蒸馏任务。
与已有工作不同的是,本文的独特切入角度是「以结果为中心」(outcome-centric)。此前工作把双层优化的内层训练视为需要绕开的负担,用特征分布或过程信号做代理;本文则直接追问:能否蒸馏出「对最终模型的影响等同于全量数据」的合成集?为此作者没有沿用需要逆 Hessian 的影响函数,而是通过展开 SGD 优化动力学并做一阶泰勒近似,把影响估计变成线性时间、可微分的量,并给出最坏情况多项式级(而非以往指数级)的误差上界(定理 2)。配合影响可加性恒等式,最小化 $\|I^-_D + I^+_S\|$ 就等价于对齐最终参数,无需真的在合成集上重训练。这使「影响匹配」首次成为可实际执行的蒸馏目标,弥合了过程对齐与结果对齐之间的优化鸿沟。
核心方法
直觉上,衡量合成数据好坏的标准不是训练过程是否相似,而是收敛后的模型是否相同:若把全量数据 $D$ 移除、再加入合成集 $S$ 后模型最终参数几乎不动,就说明 $S$ 完全替代了 $D$。Inf-Match 沿此构建目标:先在 $D$ 上用 SGD 训练 $T$ 步并记录检查点序列 $\{(\theta^t_D, \eta^t)\}$;把移除 $D$ 与添加 $S$ 造成的参数变化用一阶泰勒展开近似,写成检查点上 $H^t_D G^t_Z$、$H^t_Z G^t_D$ 形式的 Hessian-梯度乘积加权和(定理 1),该项用式 (4) 的有限差分技巧以 $\mathcal{O}(p)$ 复杂度估计,免显式逆 Hessian。蒸馏目标 $J(S)$(式 7)中,每步随机采样小批量 $B_S \subset S$、$B_D \subset D$ 并采 $m$ 个检查点近似对全部 $T$ 步求和;用 SGD-M 联合更新合成图像(学习率 50.0)与软标签(学习率 7.0),批大小 50,实验独立重复 10 次。
核心创新是结果对齐的严格形式化。利用影响函数的可加性,恒等式 $I^-_D + I^+_S \equiv \theta^*_D + I^-_D + I^+_S - \theta^*_D$ 表明:最小化 $\|I^-_D + I^+_S\|$ 等价于最小化「从 $\theta^*_D$ 出发、移除 $D$ 再添加 $S$ 后」的参数位移(Remark 1)——也就是说,不需要真的在 $S$ 上重新训练,就能通过影响估计直接衡量「用 $S$ 训练得到的模型离 $\theta^*_D$ 有多远」。这与 GM(逐点对齐梯度)和 MTT/DATM(对齐轨迹片段)有本质区别:过程匹配只是启发式代理,过程一致并不蕴含结果一致;Inf-Match 直接优化原始双层问题真正关心的量——最终参数偏移,把「合成数据应具有与真实数据相同的影响力」变成了一个可微分、可高效估计的目标。
方法步骤详情
方法流程(算法 1)分四步。第一步,轨迹采集:在 $D$ 上以 SGD 训练基础网络 $T$ 步,记录每步参数与学习率 $\{(\theta^t_D, \eta^t)\}_{t=1}^T$,作为影响估计的展开点。第二步,初始化:按 IPC 设置从 $D$ 采样真实图像构成 $S$,用最终模型给出软标签 $\hat{y}_i = f(x_i; \theta^T_D)$,图像与软标签均为可学习变量。第三步,迭代优化:每轮随机采样小批量 $B_S \subset S$ 与 $B_D \subset D$;按 DATM 式难度调度采样 $m$ 个检查点(早期偏向早期检查点学基础模式,后期偏向后期检查点编码细粒度结构);在这些检查点上计算式 (7) 的损失 $J(S)$,其中 $I^-_D$ 项保证「去掉 $D$ 再加 $S$」后参数不动,各项 Hessian-梯度乘积用式 (4) 的有限差分 HVP 以 $\mathcal{O}(p)$ 复杂度估计。第四步,更新输出:用 SGD-M 同时更新合成图像(学习率 50.0)与软标签(学习率 7.0),批大小 50,收敛后输出 $S$。
技术新颖性
技术新颖性体现在三点。其一,估计器本身:以往影响估计(Koh & Liang 的逆 Hessian 法、TracIn、DataInf 等)要么假设凸损失,要么计算昂贵;本文通过展开 SGD 动态加一阶泰勒近似,得到完全可微、线性时间的样本级估计器,适用于深度非凸网络。其二,理论保障:定理 2 给出 $|\tilde{I} - I| \le 2T^3 \ell (T+1) \eta_{\max} g + \frac{|Z|}{|D|} T^2 g$ 的最坏情形上界,误差随训练步数 $T$ 多项式增长,优于先前估计器的指数增长;界由梯度 Lipschitz 常数 $\ell$ 与梯度范数上界 $g$ 控制,意味着训练越稳定估计越准。其三,任务形式化:首次把蒸馏目标定义为影响匹配(结果对齐),并证明框架可扩展到视觉-语言蒸馏(可训练 ViT 视觉编码器 + 冻结 BERT 文本编码器 + 可训练投影层),超越 BTM、DATM 等过程匹配方法。
实验结果
分类基准上 Inf-Match 在所有 IPC 设置均第一:CIFAR-10 于 IPC=1/10/50 达 49.9%、72.5%、78.1%(IPC=50 超 NCFM 约 0.7%);CIFAR-100 达 49.3%(IPC=10)与 57.4%(IPC=50,超 NCFM 2.7%);Tiny-ImageNet 提升最大,IPC=10/50 为 31.5%、33.8%,比 NCFM 高约 4.7% 和 4.2%。跨架构实验中合成数据迁移到 ResNet-18、VGG、AlexNet 仍稳定超过 DATM(本文 45.4%–57.4%,ConvNet 上 57.4%),且蒸馏方法整体远优于 Random 与 Herding 选样。Flickr30K 视觉-语言蒸馏:200 对样本图生文 Recall@1 达 7.4%(次优 DATM 仅 1.3%);500 对文生图 14.6% vs DATM 14.1%;1000 对 16.4% 居首;200–1000 对平均比 NCFM 高 2.5%。消融(CIFAR-100, IPC=50):基线 52.2%,加真实数据初始化 53.7%,再加采样调度 55.0%(可学习标签单独 54.6%),三项齐备 57.4%,优于 DATM 55.0% 与 NCFM 54.7%。训练可视化显示收敛比 MTT 慢但终点更优,且图像逼真度与最终性能不相关;特征空间中本文样本覆盖高密度区与稀疏边缘,DM 过度集中。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| CIFAR-10 图像分类(IPC=50) | 测试准确率 (%) | 78.1 | NCFM ≈77.4 | +0.7 |
| CIFAR-100 图像分类(IPC=50) | 测试准确率 (%) | 57.4 | NCFM 54.7 / DATM 55.0 | +2.7 (vs NCFM) / +2.4 (vs DATM) |
| Tiny-ImageNet 图像分类(IPC=10) | 测试准确率 (%) | 31.5 | NCFM 26.8 | +4.7 |
| Flickr30K 图生文检索(200 对合成样本) | I2T Recall@1 (%) | 7.4 | DATM 1.3 | +6.1 |
| Flickr30K 图文检索(200–1000 对平均) | T2I/I2T Recall@1 平均 | 领先所有基线 | NCFM | 平均 +2.5 |
局限与改进
作者承认的局限:定理 2 的误差界是最坏情形,随训练步数 $T$ 以三次方速度增长(还有 $\ell$、$g$、$\eta_{\max}$ 项),虽然实证中估计影响与真实影响高度相关,但长训练视野下理论上界会变松;估计依赖预先训练并存储整条 SGD 轨迹的检查点。笔者的补充观察:分类实验主要在 CIFAR 与 Tiny-ImageNet 等中小规模数据集、以 ConvNet 为默认架构完成,尚未验证 ImageNet-1k 规模的可扩展性;Flickr30K 检索的绝对数值仍然很低(最高 16.4% Recall@1);软标签由教师模型生成,可能泄露训练数据信息,存在隐私风险;图像学习率 50.0、软标签学习率 7.0 等超参数较为特殊,迁移到新任务需重新调参;实验使用 8×A100 且每组独立重复 10 次,算力门槛不低;估计器推导基于纯 SGD(实验用带动量的 SGD-M),对 Adam 等自适应优化器的适用性未讨论。
独立分析的弱点
独立分析有以下弱点。第一,两阶段流水线与检查点存储:必须先完整训练 $D$ 并保存轨迹,$T$ 步 × $p$ 参数的存储对大模型代价高,改进方向是压缩检查点(低秩/量化)或在线式估计。第二,理论假设与实际优化器错位:误差界基于纯 SGD 与全局 Lipschitz 梯度假设,而实验使用动量且深度网络并不满足全局 Lipschitz 条件,可研究针对动量/自适应优化器的更紧界。第三,目标用参数位移的 L2 范数量化,各层参数尺度差异大,大范数层可能主导匹配,可引入逐层归一化或相对位移度量。第四,检查点采样($m$ 个)与批采样引入方差,且 DATM 式难度调度是从过程匹配方法借来的启发式,与「结果对齐」理念并不完全自洽,可探索自适应或无偏的检查点采样。第五,规模验证不足:未在 ImageNet-1k 或更大骨干(ViT/LLM 级)上验证,视觉-语言也只测了 Flickr30K,检索绝对值低;向 LLM 微调数据蒸馏的推广是明显空白。第六,隐私:软标签与蒸馏图像可能记忆训练样本,应结合差分隐私或成员推断审计。
未来方向
作者明确将开源代码(github.com/hrtan/infmatch),并指出数据集蒸馏可服务于高效数据共享、快速模型适配、隐私保护学习以及持续/联邦学习——这些都是影响匹配框架的自然落点。基于本文成果可延伸的方向包括:把该估计器用于数据修剪与数据估值(作者团队此前已有移除影响相关的工作,可无缝衔接);将影响匹配扩展到视觉-语言以外的多模态任务与 LLM 指令微调蒸馏;收紧定理 2 的误差界(数据相关界或高概率界);研究检查点数量 $m$ 与采样调度的最优策略及其理论依据;与生成式蒸馏结合,在隐空间做影响匹配后再解码为样本;以及在联邦学习中利用影响匹配处理客户端数据异构与分布漂移问题。
复现评估
复现条件中等偏友好。代码承诺在 https://github.com/hrtan/infmatch 发布(写作时尚未放出);所用数据集全部公开:CIFAR-10、CIFAR-100、Tiny-ImageNet、Flickr30K;架构简单(3 个 conv-block 的 ConvNet,Tiny-ImageNet 用 4 个,128 通道 + 池化 + ReLU + 归一化 + 线性分类器),并报告了 LeNet/AlexNet/VGG11/ResNet18 的迁移结果。论文给出了较完整的关键超参数:SGD-M 动量 0.9、批大小 50、图像学习率 50.0、软标签学习率 7.0、每组实验独立重复 10 次、检查点采样遵循 DATM 式调度、视觉-语言用可训练 ViT + 冻结 BERT + 随机初始化线性投影。算力需求为 8×A100 服务器,CIFAR 规模单卡也可能跑通,但完整复现全部表格需要较多 GPU 时。Hessian-梯度乘积用 PyTorch 的有限差分即可实现(Pearlmutter 技巧),整体实现难度中等,主要工程量在轨迹检查点的存储与管理。
论文图表
并列展示三种数据比较/生成范式:(a) 特征匹配——在真实数据上训练特征提取器,让合成数据匹配提取出的特征(DM、CAFE 等);(b) 过程匹配——对齐真实数据与合成数据引导的优化路径或梯度(GM、MTT);(c) 本文的结果匹配——直接匹配两份数据训练出的最终模型参数,且该参数不是通过实际重训练获得,而是用作者提出的影响估计器估算得到。
这张图一图道尽论文的定位:Inf-Match 属于第三种范式,并强调「不重训练、用影响估计替代」,是理解全文动机与方法差异的入口。
对 CIFAR-100 的 Wolf 类(IPC=10)做特征空间可视化:白色散点为合成样本的投影。DM 的合成样本过度集中于真实数据的高密度区域,而 Inf-Match 的合成样本同时覆盖高密度区域与分布边缘的稀疏区域,表示更均衡。
从特征分布角度直观解释了结果匹配为什么优于特征匹配:不是复制高密度样本,而是学习对最终模型有等效影响的信息,是理解方法有效性的机制性证据。