KerJEPA: Kernel Discrepancies for Euclidean Self-Supervised Learning
Eric ZimmermannEric ZimmermannHarley WiltzerJustin SzetoDavid Alvarez-MelisLester Mackey
Proposes a flexible family of kernel-regularized Joint-Embedding Predictive Architectures that uses closed-form limits of sliced maximum mean discrepancies to improve training stability and design flexibility in self-supervised learning.
Self-supervised learning enables artificial intelligence models to learn useful visual features directly from unlabeled data, forming the backbone of modern computer vision systems. A persistent challenge in this approach is avoiding representation collapse, where a network generates identical outputs for all inputs. While recent methods prevent collapse by regularizing representations toward standard target distributions using random one-dimensional projections, this projection technique introduces significant estimation noise and restricts regularizers to limited geometric assumptions.
The article aims to introduce a unified framework called KerJEPA that expands the design space for self-supervised regularization using broad families of kernel discrepancies. It evaluates alternative statistical distance measures, non-Gaussian target distributions, and analytic high-dimensional limits to enhance training stability and flexibility.
The authors conducted a theoretical and empirical study comparing multiple regularization objectives across varied representation dimensions and training schedules. They evaluated the framework on the standard ImageNette benchmark using a Vision Transformer backbone trained for up to 800 epochs. The study analyzed Maximum Mean Discrepancy, which measures distributional distances in high-dimensional feature spaces, and Kernel Stein Discrepancy, which evaluates alignment using only the score function of the target distribution without requiring sampling from the target prior.
The investigation produced several key findings. First, the standard projection-based regularization method was proven mathematically equivalent to high-dimensional Maximum Mean Discrepancy using a heavy-tailed kernel, allowing the exact high-dimensional limit to be computed directly without random projections. Second, unsliced and analytically sliced objectives significantly improved training stability and accelerated early convergence, outperforming finite-slice approximations across output dimensions ranging from 16 to 1024. Third, Kernel Stein Discrepancy using an Inverse Multiquadric kernel achieved the highest overall top-1 accuracy at approximately 91.90%, outperforming the standard baseline of 91.13%. Fourth, finite-slice approximations exhibited rising gradient variance as embedding dimensions increased, causing noticeable training instability unless large numbers of projection slices were computed.
These findings indicate that teams developing self-supervised foundation models can avoid the optimization noise and instability of random projections by adopting closed-form, analytically sliced, or unsliced kernel regularizers. Because exact and analytically derived regularizers converge faster in earlier training stages, they can reduce computational risk and total training time. Furthermore, the decoupling between projector heads and inference backbones suggests that downstream model performance is primarily driven by smooth Euclidean gradient dynamics rather than the exact choice between Gaussian or non-Gaussian target priors.
Practitioners should select regularization strategies based on batch size and dimensional constraints. For typical batch sizes where the number of samples is smaller than the required number of random projections, teams should deploy analytically sliced or unsliced Kernel Stein Discrepancy estimators to ensure stable optimization. When scaling to massive batch sizes across distributed clusters, engineering teams should evaluate efficiency tools such as random Fourier features or coreset approximations to balance computation costs.
The primary limitation of the study is its empirical reliance on ImageNette, a relatively small visual classification benchmark. While the theoretical derivations provide high confidence in the statistical mechanics of the proposed discrepancies, broader validation on large-scale datasets, downstream object detection, and semantic segmentation tasks is recommended before large-scale production deployment.
- Paper: Demystifying MMD GANs, Mikołaj Bińkowski et al. (2018). Its treatment of kernel-based maximum mean discrepancy gives useful grounding for KerJEPA’s sliced-MMD regularizers and their role in training objectives.
No sufficiently relevant recommendations were found.
