提出更优的极多分类对比学习泛化分析,突破稀有类别影响
A Refined Generalization Analysis for Extreme Multi-class Supervised Contrastive Representation Learning
- 基于类间风险集中度重构估计器,摆脱对类分布均匀性的依赖
- 样本复杂度仅与类别数R相关,达O(k),显著优于旧方法的ρ_min^{-1/2}
- 适用于长尾分布的极多分类场景,适合研究对比学习理论的研究者
对比表示学习在多个机器学习领域取得显著实证成功,但其理论样本复杂度仍不清晰。现有分析通常假设输入元组独立同分布,这一假设在实际中常被违反——对比元组来自有限标签数据池,导致元组间存在依赖。虽有近期工作使用U统计量分析此情形下的总体风险,但其技术要求每类风险均匀集中,使过风险界随ρ_min^{-1/2}增长,其中ρ_min为最罕见类的概率。这在极多分类场景下尤其悲观,因存在大量贡献小的尾部类别。本文贡献有二:第一,改进前人工作,证明样本复杂度与类别数R同阶,且不受类分布影响;第二,提出新估计器,捕捉风险在类间的集中特性,实现极多分类、长尾分布下的更紧界。在类分布的温和假设下,样本复杂度为O(k),其中k为每元组样本数。
原文摘要 · Abstract (English)
Contrastive Representation Learning (CRL) has achieved strong empirical success in multiple machine learning disciplines, yet its theoretical sample complexity remains poorly understood. Existing analyses usually assume that input tuples are identically and independently distributed, an assumption violated in most practical settings where contrastive tuples are constructed from a finite pool of labeled data, inducing dependencies among tuples. While one recent work analyzed this learning setting using U-Statistics to estimate the population risk, the techniques used therein require the risk of each class to concentrate uniformly, making excess risk bounds scale in the order of $ρ_{\min}^{-{1}/{2}}$ where $ρ_{\min}$ denotes the probability of the rarest class. Such a dependency can be overly pessimistic in the extreme multiclass settings where there are many tail classes which contribute minimally to the overall population risk. Our contributions are two-fold. Firstly, we improve upon the previous work and prove a bound with a sample complexity of the same order as the number of classes $R$, regardless of the distribution over classes. Furthermore, we formulate a different estimator that captures the concentration of the risk \textit{across classes}, enabling sharper bounds in extreme multi-class learning scenarios, especially where class distributions are long-tailed. Under mild assumptions on the class distributions, the resulting sample complexity is $\mathcal{O}(k)$ where $k$ is the number of samples per tuple.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。