🤖 AI Summary
In vision, embedding-space prediction (e.g., Joint Embedding Predictive Architecture, JEPA) substantially outperforms input-space reconstruction—yet current large language models (LLMs) still rely predominantly on token-level reconstruction for pretraining and fine-tuning, leading to weak generalization and susceptibility to overfitting. This work introduces LLM-JEPA, the first JEPA-based framework for language modeling: it jointly trains an encoder and a predictor in the embedding space, eliminating token-level reconstruction objectives. LLM-JEPA is architecture-agnostic, successfully integrated with Llama3, Gemma2, OpenELM, and Olmo. Empirical evaluation across diverse benchmarks—including NL-RX, GSM8K, Spider, and RottenTomatoes—demonstrates consistent superiority over standard training objectives, with marked improvements in generalization and robustness against overfitting. This study bridges a critical gap by extending JEPA to language modeling, establishing a novel, principled paradigm for LLM training grounded in predictive representation learning.
📝 Abstract
Large Language Model (LLM) pretraining, finetuning, and evaluation rely on input-space reconstruction and generative capabilities. Yet, it has been observed in vision that embedding-space training objectives, e.g., with Joint Embedding Predictive Architectures (JEPAs), are far superior to their input-space counterpart. That mismatch in how training is achieved between language and vision opens up a natural question: {em can language training methods learn a few tricks from the vision ones?} The lack of JEPA-style LLM is a testimony of the challenge in designing such objectives for language. In this work, we propose a first step in that direction where we develop LLM-JEPA, a JEPA based solution for LLMs applicable both to finetuning and pretraining. Thus far, LLM-JEPA is able to outperform the standard LLM training objectives by a significant margin across models, all while being robust to overfiting. Those findings are observed across numerous datasets (NL-RX, GSM8K, Spider, RottenTomatoes) and various models from the Llama3, OpenELM, Gemma2 and Olmo families. Code: https://github.com/rbalestr-lab/llm-jepa.