一句话定位
Brain-JEPA 把 i-jepa 的非生成式”隐空间预测”范式第一次系统搬到了 fMRI 脑动力学分析上:不重建被遮挡的原始 BOLD 信号(像 BrainLM 用 MAE 那样做),而是预测被遮挡脑区/时刻的表征。为了让这套范式在 fMRI 这种”没有自然顺序、空间上局部不连贯”的数据上work,作者提出两个定制组件——Brain Gradient Positioning(用功能连接梯度代替解剖坐标做 ROI 的位置编码)和Spatiotemporal Masking(Cross-ROI / Cross-Time / Double-Cross 三区域定向采样,而非 I-JEPA 式随机多块采样)。在 UK Biobank(4万+受试者)上预训练后,模型在人口学预测、疾病诊断/预后、人格特质预测等 8 个任务、5 个数据集上全面超过此前的 fMRI 基础模型 BrainLM,且在跨种族(仅用白人队列训练,泛化到亚洲队列)和线性探测(无需微调)场景下优势更明显。NeurIPS 2024 Spotlight。
背景与定位
fMRI 时间序列分析长期停留在”任务特定模型”阶段:BrainNetCNN(CNN)、BrainGNN(GNN)、Brain Network Transformer/BNT(transformer + Pearson 相关矩阵)、SwiFT(Swin Transformer 处理原始 4D fMRI)都是为单一下游任务训练的,无法利用海量未标注 fMRI 数据,泛化性有限。第一个 fMRI 基础模型 BrainLM(Caro et al., ICLR 2024)借用了 MAE 的思路:把 fMRI 时间序列当图像 patch 化,训练目标是重建被遮挡的 patch。论文指出这条路线的三个问题:(1) BOLD 信号信噪比低、信息密度稀疏,直接重建掩码区域会放大噪声或丢失关键的细微变化;(2) 已有文献(i-jepa 论文本身)证明 MAE 类生成式架构在线性探测等”开箱即用”评测上表现次优,BrainLM 因此需要额外接一个三层 MLP 做端到端微调才能达到最优效果;(3) BrainLM 只在白人队列上验证,缺乏与 SOTA 方法的对比,限制了临床适用性。
于是作者转向 i-jepa 提出的 Joint-Embedding Predictive Architecture:预测隐空间中目标块的表征而非重建原始输入,理论上能获得更高的语义抽象层级和更好的可扩展性(参见 JEPA 范式的整体蓝图 lecun-path-autonomous-machine-intelligence)。但 I-JEPA 为自然图像设计的两个组件不能直接搬过来:(1) I-JEPA 沿用 sin/cos 或图像像素坐标做位置编码,但 fMRI 的 ROI 分布在 3D 脑体积中没有自然”顺序”,解剖上相邻的 ROI 可能功能活动模式完全不同(功能分区与解剖分区不重合);(2) I-JEPA 的随机多块采样策略假设数据像 ImageNet 一样信息密度高,而 fMRI 样本量更小、信息更稀疏,需要更强的归纳偏置。Brain-JEPA 因此是 JEPA 范式向”生物医学时间序列”这一新模态的定制化迁移,与同期专注脑网络(而非脑动力学)分析的 BrainMass 是并行但方向不同的工作。
模型架构
Backbone:观测编码器 (observation encoder)、目标编码器 (target encoder,EMA 更新,不接收梯度) 与预测器 (predictor) 均为标准 ViT,用 FlashAttention(v1/v2)实现自注意力以降低显存和提高效率。预训练不使用 [CLS] token;下游评测阶段用目标编码器输出做平均池化得到全局 fMRI 表征。
模型规模(观测编码器):ViT-S(22M) / ViT-B(86M) / ViT-L(307M)。预测器是对应观测编码器的”窄”版本:ViT-S/ViT-B 的预测器深度为 6、嵌入维度分别为 192/384;ViT-L 的预测器深度为 12、嵌入维度 384。主结果、线性探测、消融实验均基于 ViT-B、预训练 300 epoch。受限于计算资源,论文未测试 ViT-H 及更大规模(自陈局限)。
Brain Gradient Positioning(空间位置编码):先算 ROI 间的非负亲和矩阵 A(i,j)=1−(1/π)cos⁻¹(cᵢcⱼᵀ/‖cᵢ‖‖cⱼ‖)(cᵢ为 ROI i 的功能连接特征),再用扩散图 (diffusion map) 方法求扩散算子 M_δ=D⁻¹L_δ, L_δ=D^(−1/δ)AD^(−1/δ)(δ=0.5,保留 ROI 间的全局关系),对 M_δ 做特征分解得到梯度矩阵 G=[ψ₁,…,ψ_m](m=30 维,用于主结果;消融另测了 3 维)。G 经一个可训练线性层映射为 Ĝ∈R^(n×d/2),与 sin/cos 时间位置编码 T∈R^(n×d/2) 拼接成最终位置嵌入 P=[T,Ĝ]∈R^(n×d)(d 为 ViT 嵌入维度)。这套坐标系是从群体级 UKB 预训练切分数据(80%)中用 BrainSpace 工具箱算出的,替代了 BrainLM 使用的纯解剖坐标。
Spatiotemporal Masking(掩码/目标采样策略):输入 fMRI(ROI 打乱后,每个 ROI 的时序按 p=16 个时间点切成 patch)划出观测块 x(在 ROI 维度范围 η_R^o 与时间维度范围 η_T^o 内随机采样,见下表),其余区域分为三个互不重叠的区域:Cross-ROI(α)、Cross-Time(β)、Double-Cross(γ)——分别对应”跨未见 ROI 泛化”、“跨未见时刻泛化”、“同时跨未见 ROI 与未见时刻泛化(最难)“。每种区域随机采样 K=1 个目标块。重叠采样 (overlapped sampling):α/β 目标从”观测块掩码 ∪ 该区域掩码”的并集中采样,γ 目标直接从 γ 区域掩码采样(Eq. 5),采样后若观测块与 α 目标重叠的 ROI、或与 β 目标重叠的 ROI 对应时间步会被从观测块中剔除,避免信息泄露。
训练目标(Eq. 6):预测器 g_φ 以位置嵌入 P 为条件,从观测表征 s_x 预测三类目标表征 ŝ_y^r,损失为 3K 个目标块上预测与真实(目标编码器输出)表征的平均 L2 距离:ℒ=(1/3K)Σ_r‖ŝ_y^r−s_y^r‖₂²。全程非生成式,不重建原始 BOLD 信号。
数据
预训练数据:UK Biobank (UKB),静息态 fMRI + 医疗记录,40,162 名参与者(年龄 44–83 岁),多站点采集,多波段扫描(TR≈0.735s)。80% 划为预训练集(同时用于计算群体级功能梯度),剩余 20% 留作站内下游评测(年龄、性别预测)。
外部评测数据集(共 5 个数据集、8 个任务):
- HCP-Aging:656 名健康老年参与者,用于人格特质(Neuroticism、Flanker 分数)与人口学(年龄、性别)预测。
- ADNI:189 名参与者用于 NC(正常对照)vs MCI(轻度认知障碍)分类;100 名认知正常参与者用于淀粉样蛋白阳性/阴性(Amyloid+/−,[18F]-Florbetapir PET SUVR ≥1.11 判阳性)分类。
- MACC(新加坡记忆-老龄-认知中心,亚洲队列):539 名参与者,NC vs MCI 分类——用于检验跨种族泛化(预训练队列 UKB 以白人为主)。
- 附录额外数据集:OASIS-3(MCI 转 AD 预测)、CamCAN(抑郁诊断)。
预处理:全部数据用 Schaefer-400(皮层)+ Tian-Scale III(皮层下)图谱分割为 n=450 个 ROI;对每个 ROI 按参与者做 robust scaling(减中位数、除四分位距)。默认输入尺寸为 160 个时间点 × 450 ROI。UKB/HCP-Aging 为多波段高时间分辨率(TR≈0.7s),ADNI/MACC 为单波段低分辨率(TR≈2s);为统一时间分辨率,对多波段数据按时间步长 3 做下采样,使全部数据集 TR 对齐到约 2 秒。下游微调/线性探测阶段,各数据集按 6:2:2 划分训练/验证/测试集。数据全部为真实人体 fMRI 记录,无仿真数据、无跨模态 co-training。
训练方法
预训练超参(Table 4):优化器 AdamW(β₁,β₂=0.9,0.999);学习率 warmup cosine schedule,起始 5×10⁻⁵ → 峰值 1×10⁻³ → 终值 1×10⁻⁶,warmup 40 epoch;权重衰减 cosine schedule,0.04→0.4;EMA 动量 schedule 线性,起始 0.996 → 终值 1(沿用 I-JEPA 惯例);总 batch size 配置为 “4 GPU × 8 梯度累积步 × 16 batch size”;patch size p=16;梯度向量维度 m=30;训练 300 epoch。
掩码区域范围(Table 6,{ROI 维度范围, 时间维度范围}):观测块 {(0.84,1.0),(0.84,1.0)};目标 α {(0.45,0.6),(0.2,0.6)};目标 β {(0.15,0.3),(0.0,0.4)};目标 γ {(0.15,0.3),(0.0,0.4)}(β、γ 采样范围数值相同,区别在于取自不同区域)。K=1(每种区域采 1 个目标块)。
下游微调(Table 5):AdamW,layer-wise lr decay 0.75(沿用 I-JEPA/MAE 惯例),base lr 0.001,weight decay 0.05,batch size 16,0 warmup epoch,训练 50 epoch。
下游线性探测:LARS 优化器,动量 0.9,base lr 0.01,batch size 64,训练 50 epoch,线性头前加一层 BatchNorm(沿用 I-JEPA 惯例)。
消融验证的设计选择:(1) Brain Gradient Positioning vs sin/cos 位置编码 vs 解剖位置编码——梯度定位全面占优;(2) Spatiotemporal Masking vs I-JEPA 原版随机多块采样——不仅性能更优,且收敛效率显著提升:消融版需要完整 300 epoch 才能达到的峰值性能,Brain-JEPA 分别只需 50/100/200 epoch(因任务而异);(3) 梯度维度 3-dim vs 30-dim——30-dim 全面占优(HCP-Aging 年龄 ρ:0.819→0.844;性别 ACC:76.96%→81.52%;ADNI Amyloid ACC:67.00%→71.00%,Table 13);(4) 剥离架构变量后单独验证 JEPA 框架本身的贡献(Table 12):BrainLM(MAE) ρ=0.832/ACC=74.39%/Amy=67.00% → BrainLM 加上本文的位置编码与数据处理贡献后 ρ=0.838/ACC=76.36%/Amy=70.00% → 换成 JEPA 框架(同样带上这些贡献)ρ=0.844/ACC=81.52%/Amy=71.00%,证明 JEPA 非生成式框架本身(而非仅仅是位置编码等外围改动)带来了性能提升。
Infra(训练 / 推理工程)
- 预训练硬件:4× NVIDIA A100(40GB),单机多卡。
- 并行/精度:论文未披露具体并行策略(数据并行推断)与混合精度设置;使用 FlashAttention 优化自注意力计算/显存。
- GPU-hours / 总训练时长:未披露。
- 推理 FPS / 控制频率 / 边缘部署:未披露——本文是离线脑影像分析基础模型,非实时控制场景,论文未给出推理速度或边缘硬件评测。
评测 benchmark
站内评测(UKB held-out 20%,Table 1)——对比 BrainNetCNN、BrainGNN、BNT、TFS(Trained-From-Scratch transformer)、BrainLM:
- 年龄:Brain-JEPA MSE 0.501(最优,BrainLM 0.612)、ρ 0.718(BrainLM 0.632)
- 性别:Brain-JEPA ACC 88.17%、F1 88.58%(BrainLM ACC 86.47%、F1 86.84%)
HCP-Aging(Table 2):年龄 MSE 0.298 / ρ 0.844(BrainLM 0.331/0.832);性别 ACC 81.52% / F1 84.26%(BrainLM 74.39%/77.51%);Neuroticism MSE 0.897 / ρ 0.307(BrainLM 0.942/0.231);Flanker MSE 0.972(与 BrainLM 0.971 基本持平)/ ρ 0.406(BrainLM 0.318)。
疾病诊断/预后(Table 3):ADNI NC/MCI ACC 76.84% / F1 86.32%(BrainLM 75.79%/85.66%);ADNI Amyloid+/− ACC 71.00% / F1 75.97%(BrainLM 67.00%/68.82%);MACC(亚洲队列)NC/MCI ACC 65.98% / F1 64.67%(BrainLM 61.65%/60.26%)——仅用白人 UKB 队列预训练,在亚洲队列上仍显著超过 BrainLM,是论文重点强调的跨种族泛化证据。
更多基线(附录 C.1,Table 7-9):加入 SVM/SVR、并发工作 BrainMass、CSM(文本式表征)、SwiFT(原始 4D fMRI)对比,Brain-JEPA 在多数任务上仍最优(如 HCP-Aging 性别 ACC 81.52% vs BrainMass 74.09%;ADNI Amyloid ACC 71.00% vs BrainMass 68.00%);OASIS-3 AD 转化 ACC 69.00%/F1 67.32%,CamCAN 抑郁诊断 ACC 72.73%/F1 67.45%,均为对比方法中最优或接近最优。
预训练数据规模缩放(附录 C.2,Table 10,用 25%/50%/75%/100% UKB 预训练数据):HCP-Aging 年龄 ρ 0.659→0.768→0.813→0.844;性别 ACC 68.03%→74.24%→77.42%→81.52%;ADNI NC/MCI ACC 67.89%→71.05%→74.74%→76.84%——随预训练数据量单调提升,未见饱和。
模型规模缩放(4.4 节,Figure 3):ViT-S/B/L 三档,越大越好(趋势见图,具体数值仅以图表形式呈现,正文未给出,标记为未披露)。
线性探测(4.5 节,Figure 4):Brain-JEPA 线性探测持续优于 BrainLM,且从微调到线性探测的性能衰减更小(具体数值同样仅在图表中,正文未给出)。
可解释性分析(4.7 节):按 Schaefer 7 网络划分(CN/DMN/DAN/LN/SAN/SMN/VN)计算 NC/MCI 分类任务下的网络级注意力分布,白人和亚洲队列上呈现一致模式,默认模式网络(DMN)、控制网络(CN)、显著性腹侧注意网络(SAN)、边缘网络(LN)被模型重点关注,与既有认知障碍神经科学文献一致(定性发现,无量化数值)。
创新点与影响
贡献:
- 首次将非生成式 JEPA 框架系统迁移到 fMRI 脑动力学基础模型,避免了 MAE 类方法在稀疏、低信噪比信号上直接重建的固有缺陷。
- Brain Gradient Positioning:用功能连接梯度(扩散图嵌入)构建 ROI 的功能坐标系,替代解剖坐标,解决了”fMRI 的 ROI 在空间上无自然序、功能分区与解剖分区不重合”的问题。
- Spatiotemporal Masking:Cross-ROI/Cross-Time/Double-Cross 三区域定向目标采样 + 重叠采样策略,比 I-JEPA 原版随机多块采样引入更强的归纳偏置,同时性能更优、收敛所需 epoch 数大幅减少。
- 在人口学预测、人格特质预测、疾病诊断/预后三大类任务、5 个数据集共 8 个任务上全面刷新 SOTA,并首次系统验证了 fMRI 基础模型的跨种族泛化能力(仅白人队列训练,亚洲队列 NC/MCI 分类显著超越此前最优)。
它改变了什么:证明了 JEPA 范式并不局限于自然图像/视频,只要针对目标模态的空间结构(无自然序的 ROI)和数据特性(稀疏低信噪比)设计对应的定位与掩码策略,就能在完全不同的生物医学时间序列模态上复现”隐空间预测优于像素级重建”的优势,尤其是在线性探测这类不做端到端微调的场景下差距更明显。这也为后续脑影像基础模型(如同一实验室后续发布的多模态版本 BrainHarmonix/Brain-Harmony)确立了范式基础。
论文自陈的局限(第 6 节):
- 受限于算力,未测试 ViT-H 等更大模型;
- 预训练数据仍偏单一(UKB 为主),更多样化的种族、站点、扫描协议、疾病队列有待纳入;
- 可解释性分析仍较粗(如皮层 vs 皮层下对比、显著 ROI/关键时间步识别)有待细化;
- 尚未探索与 MEG、EEG、T1 结构 MRI 等多模态脑数据的融合。
原始链接
- 论文(arXiv 2409.19407,NeurIPS 2024 Spotlight):https://arxiv.org/abs/2409.19407
- PDF:https://arxiv.org/pdf/2409.19407
- 官方代码:https://github.com/Eric-LRL/Brain-JEPA(README 注明后续多模态版本 BrainHarmonix/Brain-Harmony)
- HF / ModelScope / 项目页:均未找到官方发布
一手源存档(sources/)
- brain-jepa—github-readme — Eric-LRL/Brain-JEPA README 快照(sources/world-model/2024/brain-jepa—github-readme.md)
- 论文全文见 arXiv HTML(arxiv.org/html/2409.19407v1)/ PDF(arXiv 原文,不入 git):https://arxiv.org/abs/2409.19407