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

225 lines
6.2 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# Topic 2Ornstein-Uhlenbeck 过程与 Mehler 公式
> **前置知识:** [Topic 1Hermite 多项式](01_hermite_polynomials.md)、基础概率(条件期望)
> **目标:** 理解 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
```
所以 `ρ` 直接控制两个视图的相似程度:
- `ρ → 1``z' ≈ z`(几乎相同的视图)
- `ρ → 0``z'``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, g`Mehler 公式给出:
```
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ₐ = 1``w₀ = 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`](../lejepa-identifiability/experiments/lejepa_id/data.py:29) 中:
```python
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`](../lejepa-identifiability/experiments/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:谱分解与线性可识别性](03_spectral_identifiability.md)——把 Hermite 展开和 OU 衰减组合成完整的定理1证明