一句话定位
LLM-JEPA 给标准 LLM 训练目标加一个 JEPA(联合嵌入预测)附加项:保留原有下一词预测(NTP)损失的同时,让同一知识的两个”视图”(如自然语言描述与对应正则表达式/SQL/代码)在嵌入空间里互相预测,在 Llama-3/Gemma-2/OpenELM/OLMo-2 等四大模型家族、NL-RX/GSM8K/Spider/RottenTomatoes 等多个数据集、1B–8B 多种尺寸、全量微调与 LoRA 微调、乃至推理模型(Qwen3、DeepSeek-R1-Distill)上都稳定优于纯 NTP 微调/预训练,且更抗过拟合。
背景与定位
论文把表征学习分成两大阵营:生成式/重建式方法(GPT-3、PaLM、MAE)与重建无关的 JEPA(i-jepa、data2vec、V-JEPA revisited feature prediction),后者在视觉领域已被证明能带来更少偏置、更强的下游表征质量,但语言模型至今仍以 NTP 为绝对主流——因为 LLM 的评测和使用场景本质上要求”生成输入空间样本”,这与 JEPA 天然的 reconstruction-free 特性存在张力。论文承接 LeCun 2022《A Path Towards Autonomous Machine Intelligence》里提出的世界模型/JEPA 范式,尝试把这套视觉里验证过的目标”原样”搬进语言模型训练,而不是发明一个新架构。
与其区分的相关工作:SimCSE(用 dropout 生成的视图做对比学习,只产出句子表征、不含生成能力)、Meta 的 Large Concept Models(在句向量空间里做语言建模,依赖复杂的层级/聚类结构约束)、Sentence-BERT(BERT 预训练 + 语义相似度损失,同样牺牲生成能力)——这些方法都以放弃生成能力为代价换取更好的嵌入,而 LLM-JEPA 的核心设计原则是”JEPA 项作为附加正则,绝不牺牲 NTP 生成能力”。论文也指出(text,code)这类”两视图”设置此前已被用于纯生成式任务(NL→正则表达式、NL→SQL、issue→diff、数学题→程序归纳),但从未有人在这些数据上加 JEPA 风格的嵌入空间目标。
同期另一条把 JEPA 迁移到语言/推理的路线是 jepa-reasoner——它把”隐空间推理”和”token 化表达”拆成两个独立模型,本质是替换生成机制本身;LLM-JEPA 则完全不改变 LLM 的自回归生成方式,只是在训练时叠加一个作用于文本表征的辅助损失,推理阶段与原始 LLM 完全一致。
模型架构
LLM-JEPA 不引入新的 backbone,直接复用待训练的标准 decoder-only Transformer LLM 本身身兼 encoder Enc(·) 与 predictor Pred(·):
- Encoder:取输入序列最后一层、最后一个 token 的 hidden_state 作为该序列的 embedding(LLM probing 常规做法)。
- Predictor:tied-weight,不引入任何新参数——在输入末尾追加 k∈{0,1,2,3,4} 个可学习的
[PRED]特殊 token,取最后一个 PRED token 最终层的隐状态作为 Pred(Enc(·));k=0 时 Pred 是恒等映射。 - 度量 d(·,·):默认用余弦相似度(cosine),消融对比了 ℓ2 距离、MSE、InfoNCE(见”训练方法”)。
- 完整损失:L_LLM-JEPA = Σ_{ℓ=2}^{L} L_LLM(Text_{1:ℓ-1}, Text_ℓ) + λ×d(Pred(Enc(Text)), Enc(Code)),λ≥0 是 JEPA 项权重。附录进一步引入 γ 显式控制 NTP 项权重(L_LLM-JEPA=γ×NTP项+λ×JEPA项,固定 max(γ,λ)=1 以保持等效学习率);实验显示 γ=0(完全去掉 NTP 项)时模型只输出空文本,说明 NTP 项不可或缺,JEPA 项本质是正则而非替代项。
- 前向传递实现(工程核心变化,v1→v2):因果自注意力下无法用一次前向同时拿到 Enc(Text) 和 Enc(Code)(简单拼接会让后一视图的表征依赖前一视图,破坏”独立视图”假设)。v1 版本对 Text、Code 各单独做一次前向,加上原本的 NTP 前向共 3 次前向。v2 改用自定义 4D additive attention mask:把 Text 和 Code 打包进同一个 context window,构造一个”每块内因果、块间互斥”的负无穷 mask(block 数固定为 2),使两个视图各自因果、互不可见;由此把额外前向合并进 NTP 前向,总前向次数从 3 降到 2(GitHub 对应
--additive_mask参数)。 - “NTP 不能隐式最小化 JEPA 项”实验(S3.3):用 Llama-3.2-1B-Instruct + NL-RX-SYNTH 做对照——纯 NTP 训练下监控(不回传梯度)JEPA loss 发现其不会自动下降;同时对比纯 NTP 与 LLM-JEPA 两种训练下的 NTP loss 曲线几乎重合,说明 JEPA 项不会伤害生成能力。该实验里两种训练下的准确率分别为 L_LLM=51.95%、L_LLM-JEPA=71.10%,虽然二者 NTP loss 相近,NTP loss 本身无法解释这一准确率差距,论文将差距归因于
pred项(JEPA loss)——呼应 Balestriero & LeCun (2024) 在图像域发现的”自编码器不牺牲重建目标也能显著变强分类器”现象。
数据
所有实验均基于”存在自然两视图结构”的已公开数据集,论文未构造新的大规模语料,核心是复用 + 重新标注既有数据集用于微调/小规模预训练:
- NL-RX-SYNTH / NL-RX-TURK(Locascio et al. 2016):自然语言描述↔正则表达式,accuracy = 生成正则表达式与真值 exact match。
- GSM8K(Cobbe et al. 2021):小学数学应用题,accuracy = 最终数值答案 exact match。
- Spider(Yu et al. 2018):数据库 ID + 自然语言问题↔SQL 查询,accuracy = 生成 SQL 在对应 sqlite 数据库上执行结果与真值 exact match(GitHub 附带需解压的
spider_data.zipsqlite 库)。 - cestwc/paraphrase(HF,仅 train split):每组 5 条同义句,仅用于预训练——JEPA 目标让第 i 条预测第 (i+1) 条,把 5 个版本的表征约束进同一紧致子空间。
- rotten_tomatoes / yelp(HF):paraphrase 预训练后的下游微调评测,输出离散情感标签(Good/Bad 二类;Very Good/Good/Mediocre/Bad/Very Bad 五类),用生成前缀匹配真值标签判定正确(沿用 McCann 2018 / Raffel 2020 / Wei 2022 的文本分类-as-生成评测惯例);微调阶段不使用 JEPA 损失,专门用来验证”JEPA 只作用于预训练阶段”的收益是否能传导到下游。
- NQ-Open(Lee et al. 2019)、HellaSwag(Zellers et al. 2019,v2 新增):超越”文本→代码”这类天然强对应的两视图设定——NQ-Open 把 Text 定义为问题、Code 定义为答案片段(远短于其余数据集里的均衡长度);HellaSwag 把 Text 定义为上下文、Code 定义为正确续写(而非单个 ABCD 标签),二者共同的关键差异是 (i) Text 与 Code 都是”问题”本身的组成部分而非严格对应的另一形式 (ii) 上下文与续写关系比 NL→Regex/SQL 松散得多,用于检验 LLM-JEPA 在”弱两视图”结构下是否仍成立。
- Sim-vs-real / co-training:不适用(纯文本任务,无模拟/真实数据混配问题)。
训练方法
- 统一实验协议(v2):5 个固定随机种子 {82, 23, 37, 84, 4},每个 (model, dataset, config) 跑 5 次,报告均值±标准差,并用 paired one-tailed t-test 给出 p 值;训练 6 个 epoch(v1 版本为 4 个 epoch)。
- 两阶段超参搜索:先在 lr∈{1e-5, 2e-5, 4e-5, 8e-5} 上为纯 NTP baseline 挑出(model, dataset)下 4 epoch 后准确率最高的 lr,作为”最强 baseline”;再固定该 lr,在 (k, λ)∈{0,1,2,3,4}×{0.5,1,2,4} 网格上为 LLM-JEPA 搜专属超参。论文承认这带来显著调参成本,并观察到网格里相邻点准确率往往相近(提示存在更高效搜索算法的空间,但未给出)。在 QA/推理任务上进一步发现 λ 可推到 1024 而精度仍在上涨、未见平台期(Figure 9,HellaSwag + Llama-3.2-1B)。
- 消融:设计选择(S4.3,NL-RX-SYNTH,lr=2e-5, λ=1, k=1)——baseline NTP 57.29±5.32、LLM-JEPA(cosine) 71.46±1.34、ℓ2-norm 距离 2.22±0.07(灾难性失败,远低于 baseline)、MSE 70.64±2.05、把 [PRED] 前置而非后置(Prepend)68.07±2.57、反向预测 Code→Text 65.70±2.63、InfoNCE(τ=0.07)34.40±6.10(不仅低于 LLM-JEPA 也低于其余替代方案,标准差显著更大)。结论:cosine + Text→Code + 后置 [PRED] 的组合最优,其余替代设计大多仍优于纯 NTP baseline(ℓ2-norm 是唯一反例)。
- LoRA 微调:
--lora --lora_rank <N>,在 rank∈{32,64,128,256,512} 及全量微调上对比,LLM-JEPA 在每个 rank 上都显著优于同 rank 的 NTP LoRA,且在 rank 512(22.59% 可训练参数)时已追平全量 NTP 微调的准确率;同时 LoRA+LLM-JEPA 表现出比 NTP LoRA 更强的抗过拟合能力(继续训练更多 epoch 精度仍在提升,NTP LoRA 则出现明显过拟合)。 - 预训练实验:(1) Llama-3.2-1B-Instruct 从随机初始化权重在 NL-RX-SYNTH 上预训练(因数据规模有限,评测放宽为”生成以真值为前缀即算对”);(2) 在 cestwc/paraphrase 上预训练 4 epoch 后,在 rotten_tomatoes/yelp 上做 1 epoch 纯 NTP 微调评测下游效果——微调阶段本身不使用 JEPA 损失,专门隔离验证”JEPA 只作用于预训练阶段仍能带来统计显著的下游收益”。
- Loss Dropout(S5.2,加速手段):训练时按 batch 级别以概率 LD=α 随机丢弃 JEPA 项,丢弃时跳过额外的 Enc(Text)/Enc(Code) 前向;若 LD=α,则每 epoch 计算成本降为标准微调的 (2−α)× (对应 GitHub
--jepa_ratio参数,设为 1−α)。实验发现 LLM-JEPA 能容忍激进的 dropout 率(LD=0.5 或 0.75),在同等计算预算下反而比不丢弃(LD=0)取得更高准确率;并给出经验准则——保持 λ×(1−α) 近似恒定可作为联合调 λ 与 α 的实用指南。
Infra(训练 / 推理工程)
- 训练:论文正文未披露具体 GPU 型号、卡数、总 GPU-hours、并行策略与训练精度(bf16/fp16)。GitHub 提供
finetune8bh200.py/run8bh200.sh脚本用于在 NVIDIA H200 上训练规模达 8B 参数的模型,但未给出对应的卡数或耗时。 - 计算开销(论文明确量化的核心数字):v1 实现需 3 次前向(1 次 NTP + 2 次视图编码);v2 自定义 attention mask 实现降为 2 次前向;Loss Dropout 进一步把平均每步开销降到 (2−α)×(LD=0.75 时约为标准微调的 1.25×)。Table 6(loss dropout 结果表)以 4.83 PFLOPs 为一档比较不同 (LD, λ) 组合在同等计算预算下的准确率。
- 推理:LLM-JEPA 只在训练阶段引入额外前向和 [PRED] token,推理时与原始 LLM 完全一致,不引入任何额外延迟或架构改动(论文在结论中明确强调这点)——FPS/控制频率/边端硬件不适用于本工作(非在线控制或生成式部署场景)。
- 未披露:GPU 具体型号(除 8B 脚本提及 H200)、卡数、总训练时长、显存占用、混合精度设置。
评测 benchmark
以下数字均取自论文一手表格(均为 5 次运行的均值±标准差,config 列为对应最优超参):
跨模型家族(NL-RX-SYNTH 微调,Table 12):
- Llama-3.2-1B-Instruct:NTP 57.29±5.32(lr=2e-5)→ LLM-JEPA 71.46±1.34(λ=1,k=1),p=1.0e-3
- gemma-2-2b-it:NTP 33.65±3.24(lr=1e-5)→ LLM-JEPA 43.12±2.61(λ=2,k=4),p=5.5e-3
- OpenELM-1_1B-Instruct:NTP 12.07±1.81(lr=8e-5)→ LLM-JEPA 25.40±2.40(λ=4,k=3),p=5.1e-4
- OLMo-2-0425-1B-Instruct:NTP 87.09±0.36(lr=8e-5)→ LLM-JEPA 87.52±0.29(λ=2,k=0),p=2.5e-3
跨数据集(Llama-3.2-1B-Instruct 微调,Table 13):
- NL-RX-TURK:NTP 22.49±1.91 → LLM-JEPA 30.94±1.13(λ=1,k=1),p=2.4e-4
- GSM8K:NTP 32.36±0.58 → LLM-JEPA 36.36±0.20(λ=0.5,k=4),p=9.6e-5
- Spider:NTP 47.52±2.44 → LLM-JEPA 50.55±2.08(λ=1,k=3),p=4.0e-3
跨模型尺寸(NL-RX-SYNTH,Table 15;Llama-3.1-8B 用 startswith 宽松指标因难以正确终止生成):
- Llama-3.2-3B-Instruct:74.55±3.58 → 77.16±3.66(λ=2,k=0),p=0.0352
- Llama-3.1-8B-Instruct:35.77±6.60 → 63.57±16.81(λ=2.0,k=0),p=0.0131
- OLMo-2-1124-7B-Instruct:87.26±0.27 → 87.75±0.33(λ=20,k=2),p=0.0345
预训练(Table 2,NL-RX-SYNTH 从随机初始化):Llama-3.2-1B-Instruct NTP 54.38±1.70 → LLM-JEPA 60.59±1.01(λ=2,k=3),p=2.94e-4。 预训练+下游微调(Table 9,paraphrase 预训练→rotten_tomatoes/yelp 微调,微调阶段不含 JEPA 损失):Rotten Tomatoes 56.57±1.66→57.76±1.33(p=7.38e-4);Yelp 26.46±0.92→27.15±0.93(p=1.00e-3)。
超越”两视图”设定(Table 4,Llama-3.2-1B-Instruct):NQ-Open 20.12±0.41→21.59±0.40(λ=1024,k=0),p=2.44e-3;HellaSwag 69.40±0.99(lr=4e-5)→70.51±1.20(λ=1,k=3),p=0.0136。 推理模型(Table 5,GSM8K):Qwen3-1.7B 44.32±0.39→45.00±0.40(λ=1,k=0),p=0.0115;DeepSeek-R1-Distill-Qwen-1.5B 13.87±1.01→15.04±0.15(λ=0.5,k=1),p=0.0396。
LoRA vs 全量微调(Table 8,NL-RX-SYNTH,lr=2e-5,λ=1,k=1):rank 32 NTP 6.09±0.55→LLM-JEPA 7.45±1.87;rank 128 34.21±2.82→48.45±3.66;rank 512 50.18±5.15→72.41±2.94;全量微调 57.29±5.32→70.42±2.36——LoRA rank 512 下的 LLM-JEPA 已达到(甚至略超)全量 NTP 微调水平。
表征结构证据(Table 10/14,Enc(Text)−Enc(Code) 差的最小二乘线性拟合误差与 top-100 奇异值均值):base model 拟合误差 3953.11、奇异值均值 310.73;NTP 微调后 3035.01 / 341.80(结构未改善甚至更乱);LLM-JEPA k=1:4.47 / 94.84;k=0:4.04 / 16.82——LLM-JEPA 把 Text→Code 的映射压缩进一个近乎线性、维度极低的子空间,与 t-SNE 可视化里”NTP 微调打乱基座模型原有结构、LLM-JEPA 诱导出清晰聚类”的现象一致。
基线对照:本工作没有与其他 LLM 训练方法(如 SimCSE、Sentence-BERT 式目标)做直接数值对比,仅与”标准 NTP 微调/预训练”做配对显著性检验;所有结果均来自作者自建的两视图数据集实验,未在通用公开榜单(如 MMLU、HumanEval)上报告。
创新点与影响
- 首个面向 LLM 的 JEPA 训练目标:把视觉领域已验证的 JEPA 范式(预测同一知识不同视图的嵌入而非重建输入)原样搬进语言模型训练,同时保留 NTP 的生成能力——此前语言模型里的”嵌入空间正则化”方法(SimCSE、Large Concept Models、Sentence-BERT)都以牺牲生成能力为代价。
- 工程可行性持续优化:从 v1 的 3 次前向(额外开销 2×)优化到 v2 自定义 attention mask 的 2 次前向(额外开销 1×),再叠加 Loss Dropout 把开销压到 1.25×(LD=0.75)左右,为把该目标推向大规模预训练铺路。
- 可解释的表征几何证据:论文用 SVD 奇异值坍缩、最小二乘线性拟合误差、t-SNE 聚类三条独立证据一致指出 LLM-JEPA 诱导出 Text↔Code 之间近似线性、低维的映射结构,为”为什么 JEPA 损失能带来准确率提升”给出了机制性解释,而非仅报告端到端数字。
- 泛化边界的系统探测:从”天然强对应”的 NL↔正则/SQL,扩展到”松散对应”的 QA(NQ-Open)/常识补全(HellaSwag),再到推理模型(Qwen3、DeepSeek-R1-Distill)上的 GSM8K,均观测到统计显著提升,说明方法不局限于最初设计的”文本-代码”场景。
- 论文自陈局限:(1) 训练阶段仍有 2× 计算开销(虽已被 Loss Dropout 缓解,尚未完全消除);(2) 引入 λ、k 两个额外超参,网格搜索成本高,论文未给出高效调参算法;(3) 方法依赖数据集天然提供”非平凡的两个视图”,作者明确指出这是当前最大局限——尚未设计出类似视觉领域数据增强的机制,使 JEPA 目标能用于不具备天然两视图结构的任意数据集;(4) 预训练实验规模仍属小规模验证性质,尚未在生产级预训练语料上验证。
原始链接
- arXiv(v2,最新版):https://arxiv.org/abs/2509.14252
- arXiv PDF:https://arxiv.org/pdf/2509.14252
- GitHub(训练/消融/FLOPs 追踪/loss dropout 全部代码):https://github.com/rbalestr-lab/llm-jepa
一手源存档(sources/)
- llm-jepa—github-readme — GitHub README 快照(additive_mask/loss dropout/FLOPs 追踪/数据集来源,fetched 2026-07-16)
- arXiv 全文已通读 v1(2025-09-11)与 v2(2025-10-07,正文新增 Background 一节、custom attention mask 实现、design-choice 消融、QA/推理模型扩展、loss dropout 章节)两版 HTML 全文,(arXiv 原文 PDF,不入 git)