- Added 2506.09985_vjepa2.pdf (19.8 MB) - Added 2511.08544_lejepa.pdf (8.5 MB) - Added 2602.11389_causal_jepa.pdf (2.9 MB) - Added 2603.19312_leworldmodel.pdf (5.5 MB)
This commit is contained in:
@@ -134,9 +134,10 @@ L_SIG = SIGReg(h(z), N(0,I)) # 高斯正则化(防坍塌)
|
||||
```
|
||||
|
||||
### SIGReg 实现原理
|
||||
- 通过随机投影(sketching)估计嵌入分布的特征函数
|
||||
- 与标准高斯的特征函数对比,计算偏差
|
||||
- 线性时间复杂度,~50 行代码
|
||||
- 通过随机切片(sliced/sketching)将嵌入投影到 `n_slices=256` 个一维方向
|
||||
- 在 `knots=17` 个积分节点上估计投影的**特征函数**(实部 `cos` + 虚部 `sin`)
|
||||
- 与标准高斯特征函数 `φ(t)=exp(-t²/2)` 对比,按梯形权重×高斯权重加权积分
|
||||
- 线性时间复杂度,~50 行代码(见 [`losses.py:SIGReg`](lejepa-identifiability/experiments/lejepa_id/losses.py:8))
|
||||
|
||||
### 关键超参数
|
||||
| 参数 | 推荐范围 | 说明 |
|
||||
|
||||
@@ -127,7 +127,11 @@ V̂*(h(z₀)) = V*(z₀) 且 â*_{1:T}(h(z₀)) = a*_{1:T}(z₀)
|
||||
## 🔬 实验验证
|
||||
|
||||
### 实验 1:正向可识别性(验证定理 1)
|
||||
- 2D 设置,4种非线性混合函数(螺旋、正弦剪切、抛物线剪切、RealNVP)
|
||||
- 2D 设置,4种非线性混合函数(见 [`mixing.py`](JEPA/lejepa-identifiability/experiments/lejepa_id/mixing.py)):
|
||||
- `spiral`(螺旋,保测度旋转微分同胚 `g(z) = R(π‖z‖)z`)
|
||||
- `banana`(香蕉/抛物线弯曲,`x₁ = z₁ + z₀²`)
|
||||
- `sinusoid`(正弦剪切,`x₀ = z₀ + sin(1.5 z₁)`)
|
||||
- `nvp`(RealNVP 风格耦合层,`make_coupling_mixing`,任意偶数维)
|
||||
- LeJEPA 在所有情况下恢复各向同性高斯结构(旋转等价)
|
||||
- 扩展到 **1024 维**:SIGReg 和 VICReg 保持 `R² > 0.999`
|
||||
|
||||
@@ -249,40 +253,53 @@ lejepa-identifiability/
|
||||
│ ├── Dirichlet.lean # 附录E(Dirichlet 能量替代证明)
|
||||
│ └── Planning.lean # 定理4(规划等价)
|
||||
└── experiments/
|
||||
└── lejepa_id/
|
||||
├── losses.py # SIGReg、白化损失、对齐损失、InfoNCE
|
||||
├── models.py # MLP 编码器、MatchedEncoder、CNN 编码器
|
||||
├── data.py # 潜变量采样、OU 增强
|
||||
├── metrics.py # R²、正交误差、近似界量化
|
||||
└── engine.py # 训练循环(warmup + cosine LR)
|
||||
├── lejepa_id/
|
||||
│ ├── mixing.py # 非线性混合(spiral/banana/sinusoid/coupling)
|
||||
│ ├── losses.py # SIGReg、白化损失、对齐损失、InfoNCE
|
||||
│ ├── models.py # MLP 编码器、MatchedEncoder、CNN 编码器
|
||||
│ ├── data.py # 潜变量采样、OU 增强
|
||||
│ ├── metrics.py # R²、正交误差、近似界量化、Procrustes
|
||||
│ ├── reacher.py # DMC Reacher 渲染与数据集
|
||||
│ └── engine.py # 训练循环(warmup + cosine LR)
|
||||
├── run.py # 2D/scaling/gennorm/grid 统一入口
|
||||
├── run_reacher.py # Reacher 像素观测入口
|
||||
├── prerender.py # 渲染 Reacher OU/轨迹帧
|
||||
├── analysis/ # 后处理绘图与表格
|
||||
└── configs/ # 实验超参数 YAML
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 核心实现:损失函数([`losses.py`](JEPA/lejepa-identifiability/experiments/lejepa_id/losses.py))
|
||||
|
||||
#### [`SIGReg`](JEPA/lejepa-identifiability/experiments/lejepa_id/losses.py:8)(Sketched Isotropic Gaussian Regularizer)
|
||||
#### [`SIGReg`](JEPA/lejepa-identifiability/experiments/lejepa_id/losses.py:8)(Sliced characteristic-function Isotropic Gaussian Regularizer)
|
||||
|
||||
> 代码 docstring 标注为 *Sliced characteristic function regularizer*(Balestriero & LeCun 2025)。
|
||||
|
||||
```python
|
||||
class SIGReg(nn.Module):
|
||||
def __init__(self, knots=17, n_slices=256, t_max=3.0):
|
||||
# 通过随机投影(sketching)估计特征函数
|
||||
# 与标准高斯 φ(t) = exp(-t²/2) 对比
|
||||
# 梯形积分节点 + 高斯加权
|
||||
t = torch.linspace(0, t_max, knots)
|
||||
self.phi = torch.exp(-t**2 / 2) # 标准高斯特征函数
|
||||
dt = t_max / (knots - 1)
|
||||
w = torch.full((knots,), 2 * dt); w[[0, -1]] = dt # 梯形权重
|
||||
self.phi = torch.exp(-t**2 / 2) # 标准高斯特征函数
|
||||
self.weights = w * torch.exp(-t**2 / 2) # 积分权重×高斯加权
|
||||
|
||||
def forward(self, h):
|
||||
# h: (V, B, N) -> scalar
|
||||
A = F.normalize(torch.randn(...), dim=0) # 随机投影方向
|
||||
flat = h.flatten(0, 1)
|
||||
A = F.normalize(torch.randn(flat.size(-1), self.n_slices), dim=0) # 随机切片方向
|
||||
xt = (flat @ A).unsqueeze(-1) * self.t
|
||||
err = (xt.cos().mean(0) - self.phi)**2 + xt.sin().mean(0)**2
|
||||
return (err @ self.weights).mean() * flat.size(0)
|
||||
```
|
||||
|
||||
**关键设计:**
|
||||
- 用**特征函数**(Fourier 变换)而非矩匹配来度量分布差异
|
||||
- 随机投影将高维问题降为一维切片,线性时间复杂度
|
||||
- `knots=17` 个积分节点,`n_slices=256` 个随机方向
|
||||
- 用**特征函数**(Fourier 变换)的实部/虚部偏差度量分布差异,而非矩匹配
|
||||
- 随机切片(slicing)将高维问题降为一维投影,线性时间复杂度
|
||||
- `knots=17` 个积分节点,`n_slices=256` 个随机方向,`t_max=3.0`
|
||||
- 积分权重 `weights = 梯形权重 × exp(-t²/2)`:对低频段加权更高,与高斯特征函数的衰减一致
|
||||
|
||||
#### [`alignment_loss`](JEPA/lejepa-identifiability/experiments/lejepa_id/losses.py:39) 与 [`whitening_loss`](JEPA/lejepa-identifiability/experiments/lejepa_id/losses.py:31)
|
||||
|
||||
@@ -311,6 +328,32 @@ loss = infonce_loss(h, sigma)
|
||||
|
||||
---
|
||||
|
||||
### 核心实现:非线性混合([`mixing.py`](JEPA/lejepa-identifiability/experiments/lejepa_id/mixing.py))
|
||||
|
||||
> 混合函数 `g` 将独立潜变量 `z` 映射到"观测"空间 `x = g(z)`,模拟世界的非线性渲染。编码器的任务是反转它(恢复到正交等价)。
|
||||
|
||||
```python
|
||||
def mix_spiral(z):
|
||||
"""g(z) = R(π‖z‖) z —— 保测度螺旋微分同胚(旋转角随半径变化)"""
|
||||
|
||||
def mix_banana(z):
|
||||
"""香蕉形:x₀ = z₀, x₁ = z₁ + z₀²"""
|
||||
|
||||
def mix_sinusoid(z):
|
||||
"""正弦剪切:x₀ = z₀ + sin(1.5 z₁), x₁ = z₁"""
|
||||
|
||||
def make_coupling_mixing(N, n_layers=4, seed=1337):
|
||||
"""RealNVP 风格耦合层,适用于任意偶数维 N(含 N=2 的 'nvp' 情形)
|
||||
交替更新:z2 += tanh(z1 @ W) / z1 += tanh(z2 @ W),W 为正交矩阵×2"""
|
||||
```
|
||||
|
||||
**设计要点:**
|
||||
- `spiral` 保测度 → 不改变体积,是对编码器最严苛的"扭曲"测试
|
||||
- `banana`/`sinusoid` 引入低阶多项式/三角非线性
|
||||
- `coupling` 提供可扩展到高维的可逆混合,与 [`MatchedEncoder`](JEPA/lejepa-identifiability/experiments/lejepa_id/models.py:17) 结构对偶
|
||||
|
||||
---
|
||||
|
||||
### 核心实现:数据生成([`data.py`](JEPA/lejepa-identifiability/experiments/lejepa_id/data.py))
|
||||
|
||||
#### [`ou_augment`](JEPA/lejepa-identifiability/experiments/lejepa_id/data.py:29)(OU 过程增强)
|
||||
@@ -337,16 +380,19 @@ def ou_augment(z, rho, n_views=2, dist="gaussian", alpha=None):
|
||||
|
||||
```python
|
||||
def compute_all_metrics(z, x, h, h_prime, rho, N):
|
||||
# 双向 R²(线性可识别性的主要指标)
|
||||
# 混合可逆性 R²(x↔z,作为上界参考)
|
||||
r2_zx, r2_xz = bidirectional_r2(z, x)
|
||||
# 双向 R²(z↔h,线性可识别性的主要指标)
|
||||
r2_zh, r2_hz = bidirectional_r2(z, h)
|
||||
|
||||
# 正交误差(衡量 h = Qz 中 Q 的正交性)
|
||||
A = W[:N].T # 线性回归系数
|
||||
orth_err = ||A^T A - I||_F
|
||||
orth_err_normalized = orth_err / √N # 维度归一化,便于跨 N 比较
|
||||
|
||||
# 近似界量化(验证定理3)
|
||||
delta = max(L_h - 2*(1-rho)*trace_cov, 0)
|
||||
D_bound = delta / (2*rho*(1-rho))
|
||||
D_bound = delta / (2*rho*(1-rho)) # spectral_gap = 2ρ(1-ρ)
|
||||
approx_bound = D_bound + (epsilon + D_bound)**2
|
||||
|
||||
# Procrustes 距离(最优正交对齐后的误差)
|
||||
@@ -354,6 +400,8 @@ def compute_all_metrics(z, x, h, h_prime, rho, N):
|
||||
procrustes_mse = ||h - z @ Q^T||²
|
||||
```
|
||||
|
||||
**辅助函数:** [`compute_recovery_metrics`](JEPA/lejepa-identifiability/experiments/lejepa_id/metrics.py:54) 用于 Reacher 等只需 `R²(z↔h)` + 正交误差的场景,支持 `suffix` 区分 OU/轨迹编码器的指标键名。
|
||||
|
||||
---
|
||||
|
||||
### 核心实现:编码器架构([`models.py`](JEPA/lejepa-identifiability/experiments/lejepa_id/models.py))
|
||||
@@ -381,12 +429,31 @@ def train_and_evaluate(encoder, mix_fn, *, N, rho, lamb, mode="lejepa", steps=20
|
||||
```
|
||||
|
||||
**训练流程:**
|
||||
1. 采样潜变量 `z ~ N(0, I_N)`
|
||||
1. 采样潜变量 `z ~ N(0, I_N)`(或 laplace/gennorm,用于消融)
|
||||
2. OU 增强得到正样本对 `(z, z')`
|
||||
3. 混合函数 `g` 映射到观测空间 `(x, x') = (g(z), g(z'))`
|
||||
4. 编码器 `h` 映射到嵌入空间
|
||||
5. 计算 `L = λ·SIGReg + (1-λ)·Alignment`
|
||||
6. AdamW 优化
|
||||
5. 按 `mode` 计算损失:
|
||||
- `lejepa`:`L = λ·SIGReg + (1-λ)·Alignment`
|
||||
- `whiten`(VICReg 风格对照):`L = λ·Whitening + (1-λ)·Alignment`
|
||||
- `infonce`:`L = InfoNCE(h, σ)`
|
||||
6. AdamW 优化(`lr=3e-3`,warmup 占前半段,后半段 cosine 衰减)
|
||||
7. 每 `log_every` 步在固定 `z_eval` 集上评估全部指标
|
||||
|
||||
---
|
||||
|
||||
### 核心实现:Reacher 像素数据([`reacher.py`](JEPA/lejepa-identifiability/experiments/lejepa_id/reacher.py))
|
||||
|
||||
> 验证定理4的物理控制实验,基于 DeepMind Control Suite 的 `reacher / hard` 任务,通过 MuJoCo(EGL 后端)渲染 64×64 像素帧。
|
||||
|
||||
| 函数 | 作用 |
|
||||
|------|------|
|
||||
| [`render_at`](JEPA/lejepa-identifiability/experiments/lejepa_id/reacher.py:18) | 设定关节角 `qpos` 与目标位置,渲染单帧 `(3,64,64)` |
|
||||
| [`generate_ou_image_pairs`](JEPA/lejepa-identifiability/experiments/lejepa_id/reacher.py:37) | 用 OU 过程采样关节角对 `(z_t, z_{t+1})` 并渲染为图像对 |
|
||||
| [`solve_ik_grid`](JEPA/lejepa-identifiability/experiments/lejepa_id/reacher.py:66) | 200×200 网格搜索逆运动学,定位指尖到目标的关节角 |
|
||||
| [`ReacherOUDataset`](JEPA/lejepa-identifiability/experiments/lejepa_id/reacher.py:82) | 预渲染 OU 图像对 + 真值潜变量(2D 关节角)的数据集 |
|
||||
|
||||
**关键点:** 潜变量是 **2D 关节角**,像素是高度非线性的"渲染混合"。OU 编码器恢复正交等价关节角 → 规划等价;RL 轨迹编码器因数据非各向同性而失效。
|
||||
|
||||
---
|
||||
|
||||
@@ -457,7 +524,37 @@ theorem gaussian_uniqueness (lc : LatentComponent) :
|
||||
|
||||
---
|
||||
|
||||
### 定理4证明链([`Planning.lean`](JEPA/lejepa-identifiability/lean/LeJEPA/Planning.lean))
|
||||
### 定理3证明链([`Approx.lean`](JEPA/lejepa-identifiability/lean/LeJEPA/Approx.lean),Proposition 4.3)
|
||||
|
||||
**核心装配定理(已验证):**
|
||||
```lean
|
||||
theorem approximate_identifiability
|
||||
(ρ δ ε W_nl M_Q_norm total_error : ℝ) ...
|
||||
(hgap : δ ≥ 2 * ρ * (1 - ρ) * W_nl) -- 谱间隙控制非线性能量
|
||||
(hpolar : M_Q_norm ≤ ε + W_nl) -- 极分解界
|
||||
(hpythag : total_error = M_Q_norm ^ 2 + W_nl) : -- 勾股分解
|
||||
total_error ≤ δ / (2*ρ*(1-ρ)) + (ε + δ/(2*ρ*(1-ρ))) ^ 2
|
||||
```
|
||||
|
||||
**验证状态表:**
|
||||
|
||||
| 组件 | 状态 |
|
||||
|------|------|
|
||||
| 谱间隙正性 `2ρ(1-ρ) > 0` | ✅ 已验证 |
|
||||
| 非线性能量界 `W_nl ≤ D` | ✅ 已验证 |
|
||||
| 极分解 `‖M−Q‖ ≤ ε+W_nl` | 公理化 |
|
||||
| 跨阶 Hermite 正交性 | 公理化 |
|
||||
| 线性偏差平方界 | ✅ 已验证 |
|
||||
| 勾股分解 | 公理化 |
|
||||
| 单调性 `(ε+t)²+t` 递增 | ✅ 已验证 |
|
||||
| 完整界装配 | ✅ 已验证 |
|
||||
| 精确恢复 `δ=ε=0 ⟹ 误差=0`(退化为定理1) | ✅ 已验证 |
|
||||
|
||||
**附加结论(已验证):** [`bound_small_perturbation`](JEPA/lejepa-identifiability/lean/LeJEPA/Approx.lean:181)——当 `ε+D ≤ 1` 时,二次项可被一阶项控制,界简化为 `≤ 2D + ε`。
|
||||
|
||||
---
|
||||
|
||||
### 定理4证明链([`Planning.lean`](JEPA/lejepa-identifiability/lean/LeJEPA/Planning.lean),Corollary)
|
||||
|
||||
**控制问题结构:**
|
||||
```lean
|
||||
|
||||
+10
-6
@@ -96,11 +96,13 @@ D = δ / (2ρ(1-ρ))
|
||||
|
||||
| 数学概念 | 代码实现 |
|
||||
|---------|---------|
|
||||
| SIGReg 正则化 | [`losses.py:SIGReg`](../lejepa-identifiability/experiments/lejepa_id/losses.py) |
|
||||
| 非线性混合 `g`(spiral/banana/sinusoid/coupling) | [`mixing.py`](../lejepa-identifiability/experiments/lejepa_id/mixing.py) |
|
||||
| SIGReg 正则化(切片特征函数) | [`losses.py:SIGReg`](../lejepa-identifiability/experiments/lejepa_id/losses.py) |
|
||||
| 对齐损失 | [`losses.py:alignment_loss`](../lejepa-identifiability/experiments/lejepa_id/losses.py) |
|
||||
| OU 增强 | [`data.py:ou_augment`](../lejepa-identifiability/experiments/lejepa_id/data.py) |
|
||||
| R²、正交误差、近似界 | [`metrics.py:compute_all_metrics`](../lejepa-identifiability/experiments/lejepa_id/metrics.py) |
|
||||
| 训练循环 | [`engine.py:train_and_evaluate`](../lejepa-identifiability/experiments/lejepa_id/engine.py) |
|
||||
| R²、正交误差、近似界、Procrustes | [`metrics.py:compute_all_metrics`](../lejepa-identifiability/experiments/lejepa_id/metrics.py) |
|
||||
| Reacher 像素渲染 / 数据集 | [`reacher.py`](../lejepa-identifiability/experiments/lejepa_id/reacher.py) |
|
||||
| 训练循环(lejepa/whiten/infonce) | [`engine.py:train_and_evaluate`](../lejepa-identifiability/experiments/lejepa_id/engine.py) |
|
||||
|
||||
---
|
||||
|
||||
@@ -108,12 +110,14 @@ D = δ / (2ρ(1-ρ))
|
||||
|
||||
| 定理 | Lean 文件 | 验证状态 |
|
||||
|------|----------|---------|
|
||||
| 定理1(Hermite 路径) | [`lean/LeJEPA/Hermite.lean`](../lejepa-identifiability/lean/LeJEPA/Hermite.lean) | ✅ 零 sorry |
|
||||
| 定理1 / Thm 4.1(Hermite 路径) | [`lean/LeJEPA/Hermite.lean`](../lejepa-identifiability/lean/LeJEPA/Hermite.lean) | ✅ 零 sorry |
|
||||
| 定理2(高斯唯一性) | [`lean/LeJEPA/Uniqueness.lean`](../lejepa-identifiability/lean/LeJEPA/Uniqueness.lean) | ✅ 零 sorry |
|
||||
| 定理3(近似界) | [`lean/LeJEPA/Approx.lean`](../lejepa-identifiability/lean/LeJEPA/Approx.lean) | ✅ 零 sorry |
|
||||
| 定理4(规划等价) | [`lean/LeJEPA/Planning.lean`](../lejepa-identifiability/lean/LeJEPA/Planning.lean) | ✅ 零 sorry |
|
||||
| 定理3 / Prop 4.3(近似界) | [`lean/LeJEPA/Approx.lean`](../lejepa-identifiability/lean/LeJEPA/Approx.lean) | ✅ 零 sorry |
|
||||
| 定理4 / Corollary(规划等价) | [`lean/LeJEPA/Planning.lean`](../lejepa-identifiability/lean/LeJEPA/Planning.lean) | ✅ 零 sorry |
|
||||
| 附录E(Dirichlet 路径) | [`lean/LeJEPA/Dirichlet.lean`](../lejepa-identifiability/lean/LeJEPA/Dirichlet.lean) | ✅ 零 sorry |
|
||||
|
||||
> 注:Lean 工程使用 Mathlib v4.28.0,零 `sorry`;公理化组件为 Mathlib 尚未提供的标准结论(Hermite 多项式基础设施、Mazur–Ulam、等权 AM–GM 等)。
|
||||
|
||||
---
|
||||
|
||||
## 💡 核心洞见(一句话总结)
|
||||
|
||||
Reference in New Issue
Block a user