通过细粒度分析信息指数,提升多指标模型的样本效率。
Learning Orthogonal Multi-Index Models: A Fine-Grained Information Exponent Analysis
- 结合二阶与高阶项,分步学习相关子空间与精确方向。
- 样本复杂度降至 $ ilde{O}(d P^{L-1})$,优于仅用低阶项的 $ ilde{O}(P d^{L-1})$。
- 适用于低秩多指标模型,尤其适合高维稀疏结构学习。
信息指数([BAGJ21])及其扩展——在高斯单指标模型中等价于链接函数赫米特展开的最低阶次(经标签变换后)——在预测在线随机梯度下降(SGD)的样本复杂度方面发挥了重要作用。本文表明,对于多指标模型,仅关注最低阶次会忽略关键结构,导致次优率。我们研究目标函数形式为 $f_*(oldsymbol{x}) = \sum_{k=1}^{P} \phi(\mathbf{v}_k^* \cdot \boldsymbol{x})$,其中 $P \ll d$,真实方向 $\{ \mathbf{v}_k^* \}_{k=1}^P$ 正交,且 $\phi$ 的信息指数为 $L$。根据信息指数理论,当 $L = 2$ 时,仅能恢复相关子空间(而非精确方向),因二阶项具有旋转不变性;当 $L > 2$ 时,仅靠在线 SGD 恢复方向需 $\tilde{O}(P d^{L-1})$ 样本。本文证明,通过同时利用二阶与高阶项,可先用二阶项学习相关子空间,再用高阶项精确定向,整体在线 SGD 的样本复杂度为 $\tilde{O}(d P^{L-1})$。
原文摘要 · Abstract (English)
The information exponent ([BAGJ21]) and its extensions -- which are equivalent to the lowest degree in the Hermite expansion of the link function (after a potential label transform) for Gaussian single-index models -- have played an important role in predicting the sample complexity of online stochastic gradient descent (SGD) in various learning tasks. In this work, we demonstrate that, for multi-index models, focusing solely on the lowest degree can miss key structural details of the model and result in suboptimal rates. Specifically, we consider the task of learning target functions of form $f_*(\mathbf{x}) = \sum_{k=1}^{P} ϕ(\mathbf{v}_k^* \cdot \mathbf{x})$, where $P \ll d$, the ground-truth directions $\{ \mathbf{v}_k^* \}_{k=1}^P$ are orthonormal, and the information exponent of $ϕ$ is $L$. Based on the theory of information exponent, when $L = 2$, only the relevant subspace (not the exact directions) can be recovered due to the rotational invariance of the second-order terms, and when $L > 2$, recovering the directions using online SGD require $\tilde{O}(P d^{L-1})$ samples. In this work, we show that by considering both second- and higher-order terms, we can first learn the relevant space using the second-order terms, and then the exact directions using the higher-order terms, and the overall sample and complexity of online SGD is $\tilde{O}( d P^{L-1} )$.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。