A mathematical theory for understanding when abstract representations emerge in neural networks
Wang, Johnston, Fusi
Recent experiments reveal that task-relevant variables are often encoded in approximately orthogonal subspaces of the neural activity space. These disentangled low-dimensional representations are observed in multiple brain areas and across different species, and are typically the result of a process of abstraction that supports simple forms of out-of-distribution generalization. The mechanisms by which such geometries emerge remain poorly understood, and the mechanisms that have been investigated are typically unsupervised (e.g., based on variational auto-encoders). Here, we show mathematically that abstract representations of latent variables are guaranteed to appear in the last hidden layer of feedforward nonlinear networks when they are trained on tasks that depend directly on these latent variables. These abstract representations reflect the structure of the desired outputs or the semantics of the input stimuli. To investigate the neural representations that emerge in these networks, we develop an analytical framework that maps the optimization over the network weights into a mean-field problem over the distribution of neural preactivations. Applying this framework to a finite-width ReLU network, we find that its hidden layer exhibits an abstract representation at all global minima of the task objective. We further extend these analyses to two broad families of activation functions and deep feedforward architectures, demonstrating that abstract representations naturally arise in all these scenarios. Together, these results provide an explanation for the widely observed abstract representations in both the brain and artificial neural networks, as well as a mathematically tractable toolkit for understanding the emergence of different kinds of representations in task-optimized, feature-learning network models.
academic
A mathematical theory for understanding when abstract representations emerge in neural networks
This paper investigates the mathematical mechanisms underlying the emergence of abstract representations in neural networks. Experimental findings reveal that task-relevant variables are typically encoded in approximately orthogonal subspaces of neural activity space, forming decoupled low-dimensional representations. While this geometric structure supports simple out-of-distribution generalization, the mechanisms of its emergence remain unclear. The authors mathematically prove that abstract representations necessarily emerge in the final hidden layer of feedforward nonlinear networks trained on tasks dependent on latent variables. To this end, the authors develop an analytical framework that maps network weight optimization to a mean-field problem over neural pre-activation distributions.
Universality of abstract representations: Neuroscience experiments demonstrate that neural activity across multiple brain regions and species exhibits abstract representations, where task-relevant variables are encoded in approximately orthogonal subspaces
Missing mechanistic understanding: Despite the widespread existence of this geometric structure, the network mechanisms underlying its emergence remain unclear
Limitations of existing approaches: Previously studied mechanisms are primarily unsupervised methods (e.g., variational autoencoders), but pure unsupervised learning of disentangled representations faces significant challenges due to identifiability issues
Theoretical guarantees: First mathematical proof that feedforward nonlinear networks necessarily produce abstract representations under multi-task supervised learning settings
Analytical framework: Develops a general analytical tool mapping network weight optimization to mean-field problems over neural pre-activation distributions
Activation function robustness: Proves that abstract representation emergence is robust to activation function choice
Architecture extensions: Extends analysis to deep networks and recurrent networks
Neuroscience insights: Provides computational explanations for abstract representations observed in biological neural networks
Copositive programming: Handles non-convex constraints in ReLU networks
Schur convexity: Analyzes unified properties across different activation functions
Perturbation analysis: Extends results through continuity arguments
This work provides important theoretical foundations for understanding representation learning in neural networks, with mathematical frameworks and insights valuable to both neuroscience and machine learning.