[论文解读] Learning to Re-weight Examples with Optimal Transport for Imbalanced Classification
本文提出了一种基于最优传输的新型重加权方法,用于处理类别不平衡分类问题,将类别重加权建模为分布近似问题,通过最小化不平衡数据与平衡数据分布之间的最优传输距离来实现。与现有方法不同,该方法将权重学习与分类器解耦,实现了在图像、文本和点云数据集上的最先进性能,并在类别不平衡和类别数量方面表现出更强的鲁棒性。
Imbalanced data pose challenges for deep learning based classification models. One of the most widely-used approaches for tackling imbalanced data is re-weighting, where training samples are associated with different weights in the loss function. Most of existing re-weighting approaches treat the example weights as the learnable parameter and optimize the weights on the meta set, entailing expensive bilevel optimization. In this paper, we propose a novel re-weighting method based on optimal transport (OT) from a distributional point of view. Specifically, we view the training set as an imbalanced distribution over its samples, which is transported by OT to a balanced distribution obtained from the meta set. The weights of the training samples are the probability mass of the imbalanced distribution and learned by minimizing the OT distance between the two distributions. Compared with existing methods, our proposed one disengages the dependence of the weight learning on the concerned classifier at each iteration. Experiments on image, text and point cloud datasets demonstrate that our proposed re-weighting method has excellent performance, achieving state-of-the-art results in many cases and providing a promising tool for addressing the imbalanced classification issue.
研究动机与目标
- 为解决深度学习中类别不平衡问题,即模型被多数类主导且在少数类上表现不佳。
- 克服现有重加权方法将权重学习与分类器梯度耦合的局限性,导致权重优化不准确。
- 提出一种将重加权建模为分布近似问题的方法,利用最优传输技术将不平衡训练数据与一个平衡的元分布对齐。
- 开发一种鲁棒、灵活且内存高效的方案,适用于多类别数据集,通过基于原型的分布实现。
提出的方法
- 该方法将不平衡训练集表示为一个离散的经验分布 P,其中每个样本被赋予一个可学习的权重作为其概率质量。
- 将平衡的元集建模为一个具有均匀概率质量的离散经验分布 Q,代表一个平衡的目标分布。
- 通过基于特征和真实标签的代价函数,最小化 P 与 Q 之间的最优传输(OT)距离来学习权重。
- 使用 OT 损失直接训练权重网络,将权重优化与分类器梯度解耦,实现稳定且独立的学习。
- 为降低内存成本,引入了一种基于原型的 OT 损失,通过类原型而非单个样本构建 Q。
- 该方法支持端到端训练(使用可学习的权重网络)以及直接基于 OT 的权重分配,且在权重学习过程中对分类器依赖性极低。
实验结果
研究问题
- RQ1最优传输能否在不依赖分类器梯度的情况下,有效用于不平衡分类中的样本权重学习?
- RQ2与现有自动重加权方法相比,基于 OT 的重加权方法在不同数据模态下的性能和鲁棒性如何?
- RQ3当类别数量增加时,特别是长尾设置下,该方法是否仍能保持高性能?
- RQ4基于原型的 OT 变体与全样本 OT 在内存效率和准确性方面相比如何?
- RQ5该方法是否在不同数据类型(如图像、文本和 3D 点云)上具有泛化能力?
主要发现
- 在 SST-2 和 SST-5 文本分类基准上,该方法在极端不平衡(1000:100)条件下分别实现了 87.08% 和 87.13% 的准确率,优于所有基线方法,包括 Logit Adjustment 和 Hu 等人的方法。
- 在 SST-5 上,500:50 的类别不平衡条件下,该方法实现了 44.95% 的准确率,显著优于次佳方法(基于约束的重加权方法为 44.62%),表现出对高不平衡度的鲁棒性。
- 在图像分类任务中,该方法在多个数据集上均取得了最先进结果,显著优于现有的重加权和元学习基线方法。
- 基于原型的 OT 变体在性能上与全样本 OT 相当,但内存使用量大幅降低,实现了对大规模类别数据集的可扩展性。
- 该方法在图像、文本和 3D 点云数据上均表现出一致的性能,验证了其在不同数据模态下的泛化能力和鲁棒性。
- 消融研究证实,将权重学习与分类器解耦可实现更稳定、更准确的权重优化,尤其在长尾场景下表现更优。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。