[论文解读] RankingMatch: Delving into Semi-Supervised Learning with Consistency Regularization and Ranking Loss
RankingMatch 提出了一种新颖的半监督学习方法,通过引入基于排序的损失来增强一致性正则化,以鼓励同一类别的图像产生相似的模型输出。通过引入计算高效的 BatchMean 三元组损失并直接应用于模型 logit,RankingMatch 在仅使用 250 个标签的情况下实现了 CIFAR-10 上 95.13% 的准确率,在仅使用 1000 个标签的情况下在 SVHN 上达到了 97.77% 的准确率,性能达到当前最先进水平。
Semi-supervised learning (SSL) has played an important role in leveraging unlabeled data when labeled data is limited. One of the most successful SSL approaches is based on consistency regularization, which encourages the model to produce unchanged with perturbed input. However, there has been less attention spent on inputs that have the same label. Motivated by the observation that the inputs having the same label should have the similar model outputs, we propose a novel method, RankingMatch, that considers not only the perturbed inputs but also the similarity among the inputs having the same label. We especially introduce a new objective function, dubbed BatchMean Triplet loss, which has the advantage of computational efficiency while taking into account all input samples. Our RankingMatch achieves state-of-the-art performance across many standard SSL benchmarks with a variety of labeled data amounts, including 95.13% accuracy on CIFAR-10 with 250 labels, 77.65% accuracy on CIFAR-100 with 10000 labels, 97.76% accuracy on SVHN with 250 labels, and 97.77% accuracy on SVHN with 1000 labels. We also perform an ablation study to prove the efficacy of the proposed BatchMean Triplet loss against existing versions of Triplet loss.
研究动机与目标
- 为解决现有一致性正则化方法在半监督学习中的局限性,这些方法仅关注同一输入的扰动版本,而忽略了同一类别内样本之间的相互相似性。
- 通过强制同一类别下的图像产生相似的模型输出,即使它们并非彼此的扰动,从而提升模型泛化能力。
- 设计一种计算高效的三元组损失变体,该方法考虑批量中的所有样本,同时避免 BatchAll 的高成本或 BatchHard 的复杂性。
- 通过直接在模型 logit 上应用排序损失而非特征表示,统一半监督学习与度量学习。
- 通过实证验证所提出损失组件的有效性及其对不同基准和标签预算下收敛性和性能的影响。
提出的方法
- 提出 RankingMatch,一种结合交叉熵损失、一致性正则化(如 FixMatch 所示)以及新型排序损失组件的半监督学习框架。
- 将三元组损失和对比损失直接应用于模型最终的 logit(分类得分),而非学习到的特征嵌入,以在输出层面强制实现类内相似性。
- 提出 BatchMean 三元组损失作为新变体,通过在整个批量中计算最困难的正样本对和负样本对,实现计算效率与代表性之间的平衡。
- 对 logit 进行 L2 归一化以稳定训练并提升收敛性,尤其在使用 BatchMean 损失时效果显著。
- 采用双增强策略:弱增强输入生成强增强对应样本的伪标签,以强制实现一致性。
- 使排序损失具备可微性和可扩展性,避免了 BatchAll 的内存和时间开销,同时在准确率上优于 BatchHard。
实验结果
研究问题
- RQ1在仅依赖一致性正则化的基础上,强制同一类别样本的模型输出保持相似,是否能进一步提升半监督学习性能?
- RQ2直接在 logit(而非特征)上应用排序损失,是否能带来更好的泛化能力和训练稳定性?
- RQ3能否设计一种新的三元组损失变体,使其在保持 BatchAll 的代表性同时,达到 BatchHard 的计算效率?
- RQ4L2 归一化对基于排序的半监督学习方法的训练动态和最终准确率有何影响?
- RQ5在准确率和计算成本方面,所提出的 BatchMean 三元组损失相较于现有变体(BatchAll、BatchHard)的相对贡献如何?
主要发现
- 在仅使用 250 个标注样本的情况下,RankingMatch 在 CIFAR-10 上实现了 95.13% 的 top-1 准确率,优于以往最先进方法。
- 在 SVHN 上,RankingMatch 在使用 1000 个标签时达到 97.77% 的准确率,表明其在极低标签预算下仍具有强大性能。
- 在 CIFAR-10 上使用 250 个标签时,与 BatchHard 和 BatchAll 变体相比,所提出的 BatchMean 三元组损失分别将错误率降低了 50% 和 75%。
- 若不使用 L2 归一化,BatchMean 三元组损失会导致训练发散,这凸显了其在模型稳定性中的关键作用。
- BatchMean 三元组损失显著优于 BatchAll,使 SVHN 上每轮训练的平均时间减少超过 125 秒,同时将 GPU 显存使用量减半。
- t-SNE 可视化显示,与 FixMatch 和 MixMatch 相比,RankingMatch 在 logit 空间中生成了更紧凑且更清晰分离的类别簇,表明其决策边界更优。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。