Gradient descent inference in empirical risk minimization
研究了梯度下降迭代在经验风险最小化中的联合分布性质,提出去偏梯度下降推断框架,在高维均值场机制下适用于非凸损失和非高斯数据,并给出每步迭代的统计推断和泛化误差估计。
Gradient descent is one of the most widely used iterative algorithms in modern statistical learning. However, its precise algorithmic dynamics in high-dimensional settings remain only partially understood, which has limited its broader potential for statistical inference applications. This paper provides a precise, nonasymptotic joint distributional characterization of gradient descent iterates and their debiased statistics in a broad class of empirical risk minimization problems, in the so-called mean-field regime where the sample size is proportional to the signal dimension. Our nonasymptotic state evolution theory holds for both general nonconvex loss functions and non-Gaussian data, and reveals the central role of two Onsager correction matrices that precisely characterize the nontrivial dependence among all gradient descent iterates in the mean-field regime. Leveraging the joint state evolution characterization, we show that the gradient descent iterate retrieves approximate normality after a debiasing correction via a linear combination of observable loss derivative directions from all past iterates. Crucially, the debiasing coefficients are directly linked to the Onsager correction matrices, which can be estimated in a fully data-driven manner via the proposed gradient descent inference algorithm. This leads to a new algorithmic statistical inference framework based on debiased gradient descent, which (i) applies to a broad class of models with both convex and nonconvex losses, (ii) remains valid at each iteration without requiring algorithmic convergence and (iii) exhibits a certain robustness to possible model misspecification. As a by-product, our framework also provides algorithmic estimates of the generalization error at each iteration. We demonstrate our theory and inference methods in the canonical single-index regression model and a generalized logistic regression model, where the natural loss functions may exhibit arbitrarily nonconvex landscapes. Our analysis further shows that, in linear regression with squared loss, the proposed debiased gradient descent iterate eventually coincides with the debiased convex regularized estimator in a mean-field distributional sense, and the quality of statistical inference for the unknown signal aligns exactly with the generalization error achieved along the algorithmic trajectory.