From 3df832ac6b72ade87070865fd63d090dc6375170 Mon Sep 17 00:00:00 2001 From: gaojie Date: Tue, 2 Jun 2026 04:27:06 +0800 Subject: [PATCH] Add new research papers to JEPA repository - 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) --- JEPA/lejepa_resources.md | 7 +- JEPA/lejepa_world_model_notes.md | 139 +++++++++++++++--- JEPA/math/README.md | 16 +- JEPA/{ => papers}/2105.04906_vicreg.pdf | Bin JEPA/{ => papers}/2506.09985_vjepa2.pdf | 0 JEPA/{ => papers}/2511.08544_lejepa.pdf | Bin JEPA/{ => papers}/2602.11389_causal_jepa.pdf | 0 JEPA/{ => papers}/2603.19312_leworldmodel.pdf | Bin 8 files changed, 132 insertions(+), 30 deletions(-) rename JEPA/{ => papers}/2105.04906_vicreg.pdf (100%) rename JEPA/{ => papers}/2506.09985_vjepa2.pdf (100%) rename JEPA/{ => papers}/2511.08544_lejepa.pdf (100%) rename JEPA/{ => papers}/2602.11389_causal_jepa.pdf (100%) rename JEPA/{ => papers}/2603.19312_leworldmodel.pdf (100%) diff --git a/JEPA/lejepa_resources.md b/JEPA/lejepa_resources.md index b0e9f02..5db03a4 100644 --- a/JEPA/lejepa_resources.md +++ b/JEPA/lejepa_resources.md @@ -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)) ### 关键超参数 | 参数 | 推荐范围 | 说明 | diff --git a/JEPA/lejepa_world_model_notes.md b/JEPA/lejepa_world_model_notes.md index 418f799..f9c8ce2 100644 --- a/JEPA/lejepa_world_model_notes.md +++ b/JEPA/lejepa_world_model_notes.md @@ -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 diff --git a/JEPA/math/README.md b/JEPA/math/README.md index a907ac2..61b0a84 100644 --- a/JEPA/math/README.md +++ b/JEPA/math/README.md @@ -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 等)。 + --- ## 💡 核心洞见(一句话总结) diff --git a/JEPA/2105.04906_vicreg.pdf b/JEPA/papers/2105.04906_vicreg.pdf similarity index 100% rename from JEPA/2105.04906_vicreg.pdf rename to JEPA/papers/2105.04906_vicreg.pdf diff --git a/JEPA/2506.09985_vjepa2.pdf b/JEPA/papers/2506.09985_vjepa2.pdf similarity index 100% rename from JEPA/2506.09985_vjepa2.pdf rename to JEPA/papers/2506.09985_vjepa2.pdf diff --git a/JEPA/2511.08544_lejepa.pdf b/JEPA/papers/2511.08544_lejepa.pdf similarity index 100% rename from JEPA/2511.08544_lejepa.pdf rename to JEPA/papers/2511.08544_lejepa.pdf diff --git a/JEPA/2602.11389_causal_jepa.pdf b/JEPA/papers/2602.11389_causal_jepa.pdf similarity index 100% rename from JEPA/2602.11389_causal_jepa.pdf rename to JEPA/papers/2602.11389_causal_jepa.pdf diff --git a/JEPA/2603.19312_leworldmodel.pdf b/JEPA/papers/2603.19312_leworldmodel.pdf similarity index 100% rename from JEPA/2603.19312_leworldmodel.pdf rename to JEPA/papers/2603.19312_leworldmodel.pdf