Files
worldmodel/JEPA/math/02_ou_process_mehler.md
T

6.2 KiB
Raw Blame History

Topic 2Ornstein-Uhlenbeck 过程与 Mehler 公式

前置知识: Topic 1Hermite 多项式、基础概率(条件期望) 目标: 理解 LeJEPA 中"正样本对"的生成机制,以及为什么 OU 过程对高阶成分衰减更快


🎯 核心问题

LeJEPA 训练时需要"正样本对"——同一内容的两个视图 (z, z')。这对视图是怎么生成的?为什么这种生成方式会导致高阶 Hermite 成分被更强地惩罚?


🌊 什么是 Ornstein-UhlenbeckOU)过程?

物理直觉:弹簧上的粒子

想象一个粒子被弹簧拴在原点,同时受到随机扰动:

  • 弹簧力:把粒子拉回原点(均值回归)
  • 随机扰动:布朗运动噪声

这就是 OU 过程的物理图像。

数学定义(连续时间)

dz_t = -θ z_t dt + σ dW_t

其中:

  • θ > 0:均值回归速率
  • σ:噪声强度
  • W_t:标准布朗运动

LeJEPA 中的离散版本

论文使用的是离散时间 OU 过程,一步转移:

z' = ρz + √(1-ρ²) η,   η ~ N(0, I_n)

其中 ρ ∈ (0, 1)相关系数(对应连续时间的 e^{-θΔt})。


🔑 OU 过程的三个关键性质

性质 1:平稳性(Stationarity

如果 z ~ N(0, I_n),那么 z' ~ N(0, I_n)

验证:

E[z'] = ρ·E[z] + √(1-ρ²)·E[η] = 0 + 0 = 0  ✓
Var(z') = ρ²·Var(z) + (1-ρ²)·Var(η) = ρ² + (1-ρ²) = 1  ✓

意义: 正样本对 (z, z') 的边际分布相同,满足论文的"平稳性假设"。

性质 2:相关性可控

Cov(z', z) = E[z'z^T] = ρ·E[zz^T] = ρ·I_n

所以 ρ 直接控制两个视图的相似程度:

  • ρ → 1z' ≈ z(几乎相同的视图)
  • ρ → 0z'z 独立(完全不同的视图)
  • 实践中取 ρ ∈ [0.8, 0.95]

性质 3:加性噪声(Additive Noise

转移可以写成 z' = m(z) + η,其中 m(z) = ρz 是线性漂移,η 是独立噪声。这满足论文的"加性噪声假设"。


📐 Mehler 公式:OU 过程的谱定理

什么是 Mehler 公式?

Mehler 公式描述了 OU 过程的转移核transition kernel)在 Hermite 多项式基下的展开:

p(z'|z) = φ(z') · Σ_{d=0}^{∞} ρᵈ · Heₐ(z) · Heₐ(z') / d!

其中 φ(z') 是标准高斯密度。

更直观的形式:相关性公式

对任意函数 f, gMehler 公式给出:

E[f(z) · g(z')] = Σ_{d=0}^{∞} ρᵈ · ⟨f, Heₐ⟩ · ⟨g, Heₐ⟩ / d!

特别地,当 f = g = h_i(编码器的第 i 个分量)时:

E[h_i(z) · h_i(z')] = Σ_{d=0}^{∞} ρᵈ · wₐ

其中 wₐh_i 在 d 阶 Hermite 多项式上的谱权重。


🎯 核心推论:高阶成分被更强惩罚

推导过程

设编码器分量 h_i 的谱权重为 {wₐ}(满足 Σ wₐ = 1w₀ = 0)。

由 Mehler 公式:

corr_i := E[h_i(z') · h_i(z)] = Σ_{d=1}^{∞} wₐ · ρᵈ

现在比较这个值与 ρ

corr_i = Σ_{d=1}^{∞} wₐ · ρᵈ
       ≤ Σ_{d=1}^{∞} wₐ · ρ    (因为 ρᵈ ≤ ρ 对 d ≥ 1)
       = ρ · Σ_{d=1}^{∞} wₐ
       = ρ · 1 = ρ

结论: corr_i ≤ ρ,等号成立当且仅当 w₁ = 1(即 h_i 是纯线性的)。

为什么等号只在线性时成立?

如果存在某个 d₀ ≥ 2 使得 w_{d₀} > 0,那么:

w_{d₀} · ρ^{d₀} < w_{d₀} · ρ   (严格不等式,因为 ρ^{d₀} < ρ 对 d₀ ≥ 2

所以整个求和严格小于 ρ


📊 数值例子

ρ = 0.9,考虑三种编码器:

编码器 谱权重 相关性 corr_i ρ=0.9 的差距
纯线性 h(z) = z w₁ = 1 0.9¹ = 0.900 0(最优!)
纯二次 h(z) = z²-1 w₂ = 1 0.9² = 0.810 -0.090
纯三次 h(z) = z³-3z w₃ = 1 0.9³ = 0.729 -0.171
混合 w₁=0.5, w₂=0.5 各半 0.5×0.9 + 0.5×0.81 = 0.855 -0.045

结论: 非线性成分越多,相关性越低,对齐损失越大。


🔗 与 LeJEPA 训练目标的联系

LeJEPA 的对齐损失:

L_align = E[‖h(z') - h(z)‖²]
        = 2n - 2 Σᵢ E[h_i(z') · h_i(z)]
        = 2n - 2 Σᵢ corr_i

最小化 L_align ⟺ 最大化 Σᵢ corr_i

由 Mehler 公式,corr_i ≤ ρ,所以:

L_align ≥ 2n - 2nρ = 2(1-ρ)n

等号成立当且仅当每个 h_i 都是线性的!

这就是定理1的核心:最优编码器必须是线性的


🎨 直觉图示

ρ = 0.9 时,不同阶数的衰减:

d=1 (线性):  ρ¹ = 0.900  ████████████████████  ← 最大相关性
d=2 (二次):  ρ² = 0.810  ██████████████████
d=3 (三次):  ρ³ = 0.729  ████████████████
d=4 (四次):  ρ⁴ = 0.656  ██████████████
d=5 (五次):  ρ⁵ = 0.590  █████████████

非线性成分的相关性随阶数指数衰减!

🔧 代码实现

data.py 中:

def ou_augment(z, rho, n_views=2, dist="gaussian", alpha=None):
    """z' = ρz + √(1-ρ²)η"""
    fac = (1 - rho ** 2) ** 0.5
    D, N = z.shape
    eta = sample_latents(n_views * D, N, dist=dist, ...)
    eta = eta.reshape(n_views, D, N)
    return rho * z.unsqueeze(0) + fac * eta

实验配置(configs/2d.yaml)中 rho 的典型值为 0.9


小结

  1. OU 过程 生成正样本对 (z, z'),相关性由 ρ 控制
  2. 平稳性z, z' 有相同的高斯边际分布
  3. Mehler 公式OU 过程对 d 阶 Hermite 成分的相关性为 ρᵈ
  4. 核心不等式corr_i = Σ wₐ ρᵈ ≤ ρ,等号 ⟺ 纯线性
  5. 训练含义:最小化对齐损失 → 最大化相关性 → 编码器必须是线性的

➡️ 下一步

Topic 3:谱分解与线性可识别性——把 Hermite 展开和 OU 衰减组合成完整的定理1证明