arXiv:2410.11227stat.MLcs.LG2024-10ICML被引 3

在非独立同分布和依赖数据下,证明了多源非线性表征学习的泛化保证。

Guarantees for Nonlinear Representation Learning: Non-identical Covariates, Dependent Data, Fewer Samples

  • 通过多源任务联合学习共享非线性表征,再微调目标任务函数。
  • 样本数满足条件时,目标任务风险随任务数增加趋近于已知表征的回归表现。
  • 适用于数据分布不同、存在依赖性的实际场景,如医疗或跨域学习。

现代机器学习的广泛应用得益于从多个数据源中提取有意义特征的能力。然而,许多实际场景中数据在不同源间分布不一致,且在源内存在统计依赖,违反了现有理论研究的关键假设。为此,我们为从多个数据源学习通用非线性表示建立了统计保证,允许输入分布不同且可能具有依赖性。具体而言,研究了从函数类 $\mathcal F \times \mathcal G$ 中学习 $T+1$ 个函数 $f_\star^{(t)} \circ g_\star$ 的样本复杂度,其中 $f_\star^{(t)}$ 为任务相关的线性函数,$g_\star$ 为共享的非线性表示。使用每源 $N$ 个样本估计表示 $\hat g$,再用 $N'$ 个目标任务样本通过 $\hat g$ 拟合微调函数 $\hat f^{(0)}$。我们证明当 $N \gtrsim C_{\mathrm{dep}} (\mathrm{dim}(\mathcal F) + \mathrm{C}(\mathcal G)/T)$ 时,$\hat f^{(0)} \circ \hat g$ 在目标任务上的过失风险以 $ν_{\mathrm{div}} \big(\frac{\mathrm{dim}(\mathcal F)}{N'} + \frac{\mathrm{C}(\mathcal G)}{N T} \big)$ 速率衰减,其中 $C_{\mathrm{dep}}$ 表示数据依赖的影响,$ν_{\mathrm{div}}$ 为源与目标任务间(可估计的)任务多样性度量,$\mathrm C(\mathcal G)$ 为表示类 $\mathcal G$ 的复杂度。特别地,随着任务数 $T$ 增加,样本需求与风险界均收敛至 $r$-维回归情形(如同 $g_\star$ 已知),且依赖仅影响样本需求,风险界仍与独立同分布设置一致。

原文摘要 · Abstract (English)

A driving force behind the diverse applicability of modern machine learning is the ability to extract meaningful features across many sources. However, many practical domains involve data that are non-identically distributed across sources, and statistically dependent within its source, violating vital assumptions in existing theoretical studies. Toward addressing these issues, we establish statistical guarantees for learning general $\textit{nonlinear}$ representations from multiple data sources that admit different input distributions and possibly dependent data. Specifically, we study the sample-complexity of learning $T+1$ functions $f_\star^{(t)} \circ g_\star$ from a function class $\mathcal F \times \mathcal G$, where $f_\star^{(t)}$ are task specific linear functions and $g_\star$ is a shared nonlinear representation. A representation $\hat g$ is estimated using $N$ samples from each of $T$ source tasks, and a fine-tuning function $\hat f^{(0)}$ is fit using $N'$ samples from a target task passed through $\hat g$. We show that when $N \gtrsim C_{\mathrm{dep}} (\mathrm{dim}(\mathcal F) + \mathrm{C}(\mathcal G)/T)$, the excess risk of $\hat f^{(0)} \circ \hat g$ on the target task decays as $ν_{\mathrm{div}} \big(\frac{\mathrm{dim}(\mathcal F)}{N'} + \frac{\mathrm{C}(\mathcal G)}{N T} \big)$, where $C_{\mathrm{dep}}$ denotes the effect of data dependency, $ν_{\mathrm{div}}$ denotes an (estimatable) measure of $\textit{task-diversity}$ between the source and target tasks, and $\mathrm C(\mathcal G)$ denotes the complexity of the representation class $\mathcal G$. In particular, our analysis reveals: as the number of tasks $T$ increases, both the sample requirement and risk bound converge to that of $r$-dimensional regression as if $g_\star$ had been given, and the effect of dependency only enters the sample requirement, leaving the risk bound matching the iid setting.

表示学习非独立数据多源学习

Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。