Fast Generalized Neural Tangent Kernel Statistics via Trace Estimation
Abstract
The empirical state-space Neural Tangent Kernel (NTK) describes the local learning geometry of a finite-width neural network, but computing it explicitly quickly becomes impractical. We show that many useful NTK statistics that characterize, for example, the dimensionality of learned updates or how two models or learning rules relate, can instead be computed from matrix-free products using randomized trace estimation. Namely, we use Hutch++ to estimate the NTK trace, Frobenius norm, effective rank, and alignment. Furthermore, we show that the positive-semidefinite structure of the NTK yields one-sided estimators that require only forward- or reverse-mode automatic differentiation. We validate these estimators across MLPs, recurrent GRUs, and a natural-language Transformer with up to 410 million parameters. In the latter model, the state is a high-dimensional four-tensor. We demonstrate orders-of-magnitude speedups, with the fastest estimator depending on the ratio of parameter and state dimensions. Equipped with these estimators, we examine rich and lazy RNN training using hidden-state NTK alignment and use NTK alignment as a regularizer for data-scarce knowledge distillation. We find that this regularization can modestly improve generalization, especially in very data-scarce settings. Together, these results make state-space NTK diagnostics practical even at large scales.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.