Transformers & learning dynamics
Large neural networks can be viewed as interacting particle systems. The mean-field limits provide a framework for analyzing how they train, and how they use context.
Mean-field analysis and in-context learning
I use mean-field analysis to connect neural networks with their mean-field limits described by PDEs and gradient flows. For overparameterized ResNets, we relate gradient-descent training to a Wasserstein gradient flow. Under suitable assumptions, convergence of this flow yields guarantees that sufficiently wide finite networks can achieve small training loss.
Motivated by in-context learning of transformers in the long-context regime, I describe a distinguished token interacting with an evolving context distribution. We quantify how finite-context predictions and training gradients approach their infinite-context counterparts, treating context length as a statistical resource.
Transformer dynamics and attention scaling
Repeated attention layers can make distinct token representations increasingly similar, a phenomenon known as oversmoothing. For linear graph neural networks, I use Lyapunov exponents to quantify this effect and show how residual connections mitigate it. For a mean-field transformer model, Wasserstein gradient flows and synchronization yield quantitative clustering guarantees under suitable parameter conditions.
As context grows, attention can also become diluted. An empirical strategy used in many LLMs is to allow the temperature parameter in each self-attention layer to scale with the context length. Using order statistics and extreme-value theory, we identify critical temperature scalings and characterize different scaling regimes in an attention row, where attention spreads over many tokens, concentrates on a few, or selects a single dominant token.
For papers, preprints, and related work, see Publications. Research code is available on GitHub ↗.