Learn from your own latents and not from tokens: A sample-complexity theory
Daniel Korchinski ⋅ Alessandro Favero ⋅ Matthieu Wyart
Abstract
Generative models, from diffusion models to large language models, achieve remarkable performance but at a cost in training data orders of magnitude larger than what biological learners require. An alternative paradigm has emerged in which networks are trained to predict their \emph{own} latent representations of related views or masked regions, as in data2vec and JEPA -- an idea related to predictive-coding accounts of the cortex. Despite strong empirical results, the theoretical understanding of these methods remains limited. Central questions include: by how much does latent prediction actually improve data efficiency? Is there a benefit to stacking such methods into multi-scale hierarchies? We answer both using as data a tractable probabilistic context-free grammar that captures the compositional structure of natural language and images. We prove that latent prediction recovers the full latent tree of depth $L$ from a number of samples scaling as $m^3$, where $m$ is the number of production rules per symbol. This is much fewer than the $m^L$ samples required by supervised learning and the $m^{L+1}$ required by token-level SSL. We confirm this bound with (i) a hierarchical clustering algorithm, (ii) an end-to-end neural network whose predictor-clusterer modules predict their own latents at each level via gradient descent, and (iii) the first sample-complexity analysis of data2vec, which we show implicitly performs hierarchical latent prediction. This suggests that explicit stacking such as H-JEPA is largely redundant.
Chat is not available.
Successful Page Load