GRAIN: Group Aggregation via Min-Norm Objective for Stable Learning
Abstract
Two training runs that differ only in the random seed can produce noticeably different models. The problem is long-standing, but it is most acute in the overparameterized regime of modern deep learning, where large models fine-tuned on limited data traverse flat loss landscapes with many near-equivalent minima, and it is most costly for large pretrained models (LPMs), whose training is expensive, whose downstream data is scarce, and for which repeated runs to average out variance are prohibitive. We trace a common source of this instability to the update rule itself: when the gradients of different groups of examples conflict in direction, their arithmetic mean can be small or even zero, stalling training while the loss is still high. We introduce GRAIN (GRoup Aggregation via mIN-norm objective), a lightweight, optimizer-agnostic procedure that replaces the arithmetic mean, both within and across mini-batches, with the min-norm convex combination of group-wise gradients. The resulting update is guaranteed to conflict with no group gradient, so no group loss increases to first order. Under mild smoothness and absolute-continuity assumptions it differs almost surely from the mean, which yields a uniform-stability bound strictly tighter than SGD's. Empirically, across generation, classification, and regression at LPM scale, GRAIN is the only method in our suite with no collapsed run over 10 seeds, while improving mean performance and reducing run-to-run variance over a broad set of baselines, at no extra wall-clock or storage cost under data-parallel training.
Then back it, or bet against it.
Related papers
Open the market on this paper to see 7 more related papers.