中文

D2 剪枝:用于平衡数据剪枝中多样性与难度的消息传递

机器学习 2023-10-13 v1 人工智能 计算与语言 计算机视觉与模式识别

摘要

分析理论表明,在固定数据预算下训练模型时,更高质量的数据可带来更低的测试误差。此外,若数据集能去除冗余,则模型可在更低计算预算下训练而不损性能。核心集选择(或称数据剪枝)旨在选取训练数据的子集,以最大化在该子集上训练模型的性能,该子集亦称核心集。现有两类主导方法:(1)基于几何的数据选择以最大化核心集中的数据多样性;(2)基于训练动态为样本分配难度分数的函数。优化数据多样性导致核心集偏向较易样本,而按难度排序选择则遗漏深度学习模型训练所需的简单样本。这表明数据多样性与重要性分数是核心集选择中需联合考虑的互补因素。我们将数据集表示为无向图,并提出一种新颖剪枝算法 D2 Pruning,其通过该数据集图上的前向与反向消息传递进行核心集选择。D2 Pruning 通过纳入数据集中相邻样本的难度的难度来更新各样本的难例分数。随后,这些更新后的难度分数指导基于图的采样方法,选取涵盖数据集空间中多样且困难区域的核心集。我们在多种视觉与语言数据集上评估了方法的监督与自监督版本。结果表明,在高达 70% 剪枝率下,D2 Pruning 优于先前 SOTA 方法。此外,我们发现使用 D2 Pruning 过滤大型多模态数据集可提升数据集多样性并改善预训练模型的泛化能力。

关键词

引用

@article{arxiv.2310.07931,
  title  = {D2 Pruning: Message Passing for Balancing Diversity and Difficulty in Data Pruning},
  author = {Adyasha Maharana and Prateek Yadav and Mohit Bansal},
  journal= {arXiv preprint arXiv:2310.07931},
  year   = {2023}
}

备注

17 pages (Our code is available at https://github.com/adymaharana/d2pruning)