Low-Rank Friction for Memory-Efficient Transformer Pretraining
In the authors' words
iKFAD is a recently proposed optimiser that replaces adaptive learning rates with adaptive friction in the momentum dynamics, yet performs as well as Adam. Its limitation is that the full friction tensor carries the same memory overhead per layer as Adam's second-moment buffer. Here we replace iKFAD's friction tensor with a rank-1 outer-product factorisation built from row and column momentum statistics, resulting in Rank-1 iKFAD (R-iKFAD). This reduces the friction memory footprint from to per layer, which approximately halves iKFAD's total optimiser state. Despite this reduction, R-iKFAD maintains parity in performance with iKFAD: experiments on GPT2-Nano, TinyViT, DistilBERT and GPT2-S confirm that it matches or exceeds iKFAD while nearly halving the memory footprint and remaining comparably robust to hyperparameters. We analyse the continuous-time dynamics in two damping regimes. For linear damping () we prove exponential convergence under strong convexity. For , the preferred option in our experiments, the friction is generated entirely from past momentum and switches off as the momentum vanishes, so geometric convergence cannot be shown. We nonetheless prove convergence to the minimiser, together with matching upper and lower bounds on the energy: of order when the regularisation scale is zero, and of order when it is positive. To our knowledge this is the first convergence rate for a rank-1 factored optimiser in continuous time, and the first such result that does not require positive damping.
Appeared: Monday, September 28. arXiv. Preprint, not yet peer-reviewed.