Yann LeCun 转推 SIGReg 博客解释 JEPA 防崩溃机制
RT @reza_byt: 基于 JEPA 的世界建模最近因一种名为 "SIGReg"(由 @ylecun 和 @randall_balestr 提出)的新型防崩溃机制而受到关注。数学很干净,但很少从第一性原理解释,所以我将其分解成一篇详细的博客文章 📰
以下是摘要(🧵):
问题:在 JEPA 中,预测损失的两端都经过同一个编码器。将每个输入映射到一个固定点,预测器轻松匹配,损失恰好为零。完美得分,零信息。梯度下降默认会找到这个解。
先前的工作通过 stop-gradients、EMA 师生模型、冻结的预训练编码器和带有 6 个以上手工调整系数的 VICReg 损失来修补。每个都会增加不稳定性、超参数或对他人预训练的依赖。
---
SIGReg 用一个基于单一主张的正则化器替换了所有这些:迫使嵌入批次看起来像各向同性高斯的样本 N(0, I)。LeJEPA 证明这并非随意。它是跨线性和非线性探针最小化最坏情况下游风险的分布。
但“使嵌入高斯化”说起来容易,计算起来难。你无法从小批次中估计 200+ 维的密度。解决方法是四个经典结果的链条:用傅里叶变换代替密度比较,将比较转化为单个标量,用约 16 个点近似积分,并用 1936 年的定理将整个过程从 1D 推广到任意维。每一步都很简单。这个堆叠就是它工作的原因。阅读本线程中的后续帖子以理解每一步。
---
但为何它可证明有效:SIGReg = 0 的唯一分布是 N(0, I),它通过构造是满秩的,每个特征值等于 1。坍缩的低秩编码器不可能是最小值。不仅不太可能收敛到那里——数学上被排除。
训练循环保持简单:
total loss = prediction + λ·SIGReg
编码器从两项获得梯度;预测器只从预测项获得梯度。无需交替更新、无需 stop-gradients、无需双时间尺度技巧。
现在查看下方的分解。👇
以下是摘要(🧵):
问题:在 JEPA 中,预测损失的两端都经过同一个编码器。将每个输入映射到一个固定点,预测器轻松匹配,损失恰好为零。完美得分,零信息。梯度下降默认会找到这个解。
先前的工作通过 stop-gradients、EMA 师生模型、冻结的预训练编码器和带有 6 个以上手工调整系数的 VICReg 损失来修补。每个都会增加不稳定性、超参数或对他人预训练的依赖。
---
SIGReg 用一个基于单一主张的正则化器替换了所有这些:迫使嵌入批次看起来像各向同性高斯的样本 N(0, I)。LeJEPA 证明这并非随意。它是跨线性和非线性探针最小化最坏情况下游风险的分布。
但“使嵌入高斯化”说起来容易,计算起来难。你无法从小批次中估计 200+ 维的密度。解决方法是四个经典结果的链条:用傅里叶变换代替密度比较,将比较转化为单个标量,用约 16 个点近似积分,并用 1936 年的定理将整个过程从 1D 推广到任意维。每一步都很简单。这个堆叠就是它工作的原因。阅读本线程中的后续帖子以理解每一步。
---
但为何它可证明有效:SIGReg = 0 的唯一分布是 N(0, I),它通过构造是满秩的,每个特征值等于 1。坍缩的低秩编码器不可能是最小值。不仅不太可能收敛到那里——数学上被排除。
训练循环保持简单:
total loss = prediction + λ·SIGReg
编码器从两项获得梯度;预测器只从预测项获得梯度。无需交替更新、无需 stop-gradients、无需双时间尺度技巧。
现在查看下方的分解。👇