225 lines
6.2 KiB
Markdown
225 lines
6.2 KiB
Markdown
# Topic 2:Ornstein-Uhlenbeck 过程与 Mehler 公式
|
||
|
||
> **前置知识:** [Topic 1:Hermite 多项式](01_hermite_polynomials.md)、基础概率(条件期望)
|
||
> **目标:** 理解 LeJEPA 中"正样本对"的生成机制,以及为什么 OU 过程对高阶成分衰减更快
|
||
|
||
---
|
||
|
||
## 🎯 核心问题
|
||
|
||
LeJEPA 训练时需要"正样本对"——同一内容的两个视图 `(z, z')`。这对视图是怎么生成的?为什么这种生成方式会导致高阶 Hermite 成分被更强地惩罚?
|
||
|
||
---
|
||
|
||
## 🌊 什么是 Ornstein-Uhlenbeck(OU)过程?
|
||
|
||
### 物理直觉:弹簧上的粒子
|
||
|
||
想象一个粒子被弹簧拴在原点,同时受到随机扰动:
|
||
- **弹簧力**:把粒子拉回原点(均值回归)
|
||
- **随机扰动**:布朗运动噪声
|
||
|
||
这就是 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证明
|