refactor: 将子模块转为普通目录,移除外部 git 依赖
Sync to site1 / sync (push) Has been cancelled

- 移除 JEPA/lejepa-identifiability 子模块 gitlink
- 移除 research/multiply/MultiPLY 子模块 gitlink
- 删除 .gitmodules(不再有外部 URL 依赖)
- 两个目录内容作为普通文件纳入主仓库追踪
- 删除各自内部 .git 目录,消除嵌套 git 仓库
This commit is contained in:
gaojie
2026-06-05 17:14:01 +08:00
parent cb629f18a1
commit c66855adfc
208 changed files with 23296 additions and 9 deletions
@@ -0,0 +1,371 @@
"""
Reacher trajectory distribution analysis figures.
Produces two figures for the paper:
1. Scatter grid: stationary marginal + per-delta 2D transition differences
and per-dim (z_t, z_{t+delta}) scatters, annotated with R² and rho.
2. rho-vs-SIGReg scatter: three panels (z_0, z_1, joint), colored by R²,
showing the dual constraint that identifiability requires both
rho off from 1 and approximately-Gaussian transition shape.
Usage:
python -m analysis.make_reacher_distributions \
--results_dir results/reacher \
--data_path data/reacher.h5 \
--out_dir figures/reacher
"""
import argparse
import json
from collections import defaultdict
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
from scipy.stats import pearsonr
DELTAS = [1, 2, 4, 8, 16, 32, 64]
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
N_MAX = 100_000
# ═════════════════════════════════════════════════════════════════════════════
# SIGReg (EppsPulley) — matches LeJEPA Algorithm 1
# ═════════════════════════════════════════════════════════════════════════════
def _sigreg_nd(x, num_slices=64, n_knots=17, seed=0, device=DEVICE):
"""EP via random slicing on (N, K)."""
x = torch.as_tensor(np.asarray(x), dtype=torch.float32, device=device)
if x.dim() == 1:
x = x[:, None]
N, K = x.shape
g = torch.Generator(device=device).manual_seed(seed)
A = torch.randn(K, num_slices, generator=g, device=device)
A = A / A.norm(p=2, dim=0)
t = torch.linspace(-5, 5, n_knots, device=device)
phi = torch.exp(-0.5 * t ** 2)
zt = (x @ A).unsqueeze(2) * t
cm, sm = torch.cos(zt).mean(0), torch.sin(zt).mean(0)
err = ((cm - phi) ** 2 + sm ** 2) * phi
return (torch.trapz(err, t, dim=1) * N).mean().item()
def _sigreg_1d(x, n_knots=17, device=DEVICE):
"""EP directly on 1D (no slicing)."""
x = torch.as_tensor(np.asarray(x).reshape(-1), dtype=torch.float32, device=device)
N = x.shape[0]
t = torch.linspace(-5, 5, n_knots, device=device)
phi = torch.exp(-0.5 * t ** 2)
zt = x.unsqueeze(1) * t
cm, sm = torch.cos(zt).mean(0), torch.sin(zt).mean(0)
err = ((cm - phi) ** 2 + sm ** 2) * phi
return (torch.trapz(err, t) * N).item()
def _zscore(x):
return (x - x.mean(0, keepdims=True)) / (x.std(0, keepdims=True) + 1e-8)
def measure(x, n_draws=20, N_max=N_MAX, num_slices=64):
"""
SIGReg raw + zscored for joint (K-d) and per-dim marginals.
Averages over n_draws random subsamples / projections.
Returns dict mapping key -> (mean, std) over draws.
"""
x = np.asarray(x)
if x.ndim == 1:
x = x[:, None]
N, K = x.shape
rng = np.random.default_rng(0)
keys = ["joint_raw", "joint_zs"]
keys += [f"marg_{k}_raw" for k in range(K)]
keys += [f"marg_{k}_zs" for k in range(K)]
buf = {k: [] for k in keys}
for s in range(n_draws):
xs = x if N <= N_max else x[rng.choice(N, N_max, replace=False)]
xz = _zscore(xs)
buf["joint_raw"].append(_sigreg_nd(xs, num_slices, seed=s))
buf["joint_zs"].append(_sigreg_nd(xz, num_slices, seed=s))
for k in range(K):
buf[f"marg_{k}_raw"].append(_sigreg_1d(xs[:, k]))
buf[f"marg_{k}_zs"].append(_sigreg_1d(xz[:, k]))
return {k: (np.mean(v), np.std(v)) for k, v in buf.items()}
# ═════════════════════════════════════════════════════════════════════════════
# R² loading — per-seed best-lambda tuning, median across seeds
# ═════════════════════════════════════════════════════════════════════════════
def get_best_r2_per_delta(results_dir, agg="median"):
"""
For each (delta, seed), pick the lambda with best R²; aggregate seeds
with median (robust to outliers) or mean.
"""
by_dls = defaultdict(list)
for p in Path(results_dir).rglob("result.json"):
r = json.load(open(p))
if "delta" not in r or r.get("rho") is not None:
continue
by_dls[(r["delta"], r["lamb"], r.get("seed", 0))].append(r)
deltas = sorted({k[0] for k in by_dls})
lambs = sorted({k[1] for k in by_dls})
seeds = sorted({k[2] for k in by_dls})
agg_fn = np.median if agg == "median" else np.mean
best = {}
for delta in deltas:
per_seed = {"r2": [], "d0": [], "d1": [], "lamb": []}
for seed in seeds:
best_lamb, best_r2 = None, -np.inf
for lamb in lambs:
runs = by_dls.get((delta, lamb, seed), [])
if not runs:
continue
r2 = np.mean([r["r2_hz"] for r in runs])
if r2 > best_r2:
best_r2, best_lamb = r2, lamb
if best_lamb is None:
continue
runs = by_dls[(delta, best_lamb, seed)]
per_seed["r2"].append(np.mean([r["r2_hz"] for r in runs]))
per_seed["d0"].append(np.mean([r["r2_hz_per_dim"][0] for r in runs]))
per_seed["d1"].append(np.mean([r["r2_hz_per_dim"][1] for r in runs]))
per_seed["lamb"].append(best_lamb)
best[delta] = {
"r2": agg_fn(per_seed["r2"]),
"r2_dim0": agg_fn(per_seed["d0"]),
"r2_dim1": agg_fn(per_seed["d1"]),
"lamb": np.median(per_seed["lamb"]),
"n_seeds": len(per_seed["r2"]),
}
return best
def best_ou_rho(results_dir):
"""Return the rho of the OU run with highest mean R² across seeds/lambdas."""
grouped = defaultdict(list)
for p in Path(results_dir).rglob("result.json"):
r = json.load(open(p))
if "rho" not in r or r.get("rho") is None:
continue
grouped[(r["rho"], r["lamb"])].append(r)
if not grouped:
return None
best_rho, best_mean = None, -np.inf
for (rho, lamb), runs in grouped.items():
m = np.mean([r["r2_hz"] for r in runs])
if m > best_mean:
best_mean, best_rho = m, rho
return best_rho
# ═════════════════════════════════════════════════════════════════════════════
# Figure 1: scatter grid
# ═════════════════════════════════════════════════════════════════════════════
def make_scatter_grid(episodes, r2_dict, save_path, sub=1, s=1e-4):
"""
Left column: stationary marginal scatter of (z_0, z_1).
Top row, remaining columns: 2D transition-difference scatter per delta.
Bottom row, remaining columns: per-dim (z_t, z_{t+delta}) scatter.
Titles show R² and rho.
"""
fig = plt.figure(figsize=0.85 * np.array((1 + 3 * len(DELTAS), 5)))
gs = fig.add_gridspec(2, 2 + len(DELTAS))
# stationary marginal (spans both rows, first two cols)
ax = fig.add_subplot(gs[:2, :2])
ax.scatter(*episodes.reshape(-1, 2)[::sub].T, s=s * 10)
ax.set_title("Marginal")
ax.grid()
ax.set_xlabel(r"$z_0$ (shoulder)")
ax.set_ylabel(r"$z_1$ (wrist)")
for i, delta in enumerate(DELTAS):
# top row: 2D transition differences
ax = fig.add_subplot(gs[0, 2 + i])
transitions = episodes[:, delta:] - episodes[:, :-delta]
ax.scatter(*transitions.reshape(-1, 2)[::sub].T, s=s)
r2_d0 = r2_dict[delta]["r2_dim0"]
r2_d1 = r2_dict[delta]["r2_dim1"]
ax.set_title(r"$\Delta=$" + f"{delta}" + "\n"
r"$R^2=(%.2f, %.2f)$" % (r2_d0, r2_d1))
ax.grid()
# bottom row: per-dim (z_t, z_{t+delta}) with rho
ax = fig.add_subplot(gs[1, 2 + i])
a = episodes[:, delta:, 0].flatten()[::sub]
b = episodes[:, :-delta, 0].flatten()[::sub]
rho0 = pearsonr(a, b)[0]
ax.scatter(a, b, s=s)
c = episodes[:, delta:, 1].flatten()[::sub]
d = episodes[:, :-delta, 1].flatten()[::sub]
rho1 = pearsonr(c, d)[0]
ax.scatter(c, d, s=s)
ax.set_title(r"$\rho=(%.2f, %.2f)$" % (rho0, rho1))
ax.grid()
if i == 0:
ax.legend([r"$z_0$ (shoulder)", r"$z_1$ (wrist)"], loc="upper left")
plt.tight_layout()
plt.savefig(save_path, dpi=500, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# Figure 2: rho vs SIGReg scatter (3 panels)
# ═════════════════════════════════════════════════════════════════════════════
def make_rho_vs_sigreg(trans, r2_dict, save_path, best_ou_rho_val=None,
gaussian_floor=1.2):
"""
Three panels (z_0, z_1, joint) showing rho vs SIGReg(zscored),
colored by R². Vertical line marks the best OU rho for reference.
"""
fig, axes = plt.subplots(1, 3, figsize=(16, 4.5))
names = [r"$z_0$ (shoulder)", r"$z_1$ (wrist)", "joint (avg across dims)"]
all_r2 = []
for d in DELTAS:
all_r2.append(r2_dict[d]["r2_dim0"])
all_r2.append(r2_dict[d]["r2_dim1"])
all_r2.append(r2_dict[d]["r2"])
vmin, vmax = min(all_r2), max(all_r2)
def _one(ax, rhos, sigs, errs, r2s, xlabel, ylabel, title):
if best_ou_rho_val is not None:
ax.axvline(best_ou_rho_val, color="crimson", lw=1.4, ls="--",
alpha=0.8, label=fr"best OU $\rho={best_ou_rho_val:.2f}$",
zorder=1)
sc = ax.scatter(rhos, sigs, c=r2s, cmap="viridis", s=140,
vmin=vmin, vmax=vmax,
edgecolors="black", linewidths=0.8, zorder=3)
ax.errorbar(rhos, sigs, yerr=errs, fmt="none", ecolor="gray",
alpha=0.5, zorder=2)
for d, r, s in zip(DELTAS, rhos, sigs):
ax.annotate(f"Δ={d}", (r, s), xytext=(6, 6),
textcoords="offset points", fontsize=9)
ax.axhline(gaussian_floor, color="red", lw=1, ls=":", alpha=0.6,
label="Gaussian floor")
ax.set_yscale("log")
ax.set_xlabel(xlabel)
ax.set_ylabel(ylabel)
ax.set_title(title)
ax.grid(alpha=0.3, which="both")
ax.legend(loc="lower left", fontsize=8)
return sc
# per-dim panels
for k in range(2):
rhos = np.array([trans[d]["rho"][k] for d in DELTAS])
sigs = np.array([trans[d]["sig"][f"marg_{k}_zs"][0] for d in DELTAS])
errs = np.array([trans[d]["sig"][f"marg_{k}_zs"][1] for d in DELTAS])
r2s = np.array([r2_dict[d][f"r2_dim{k}"] for d in DELTAS])
sc = _one(axes[k], rhos, sigs, errs, r2s,
xlabel=r"auto-correlation $\rho$",
ylabel="SIGReg (zscored, marginal)",
title=names[k])
plt.colorbar(sc, ax=axes[k], label=r"$R^2$")
# joint panel
rhos_avg = np.array([np.mean(trans[d]["rho"]) for d in DELTAS])
sigs_j = np.array([trans[d]["sig"]["joint_zs"][0] for d in DELTAS])
errs_j = np.array([trans[d]["sig"]["joint_zs"][1] for d in DELTAS])
r2s_avg = np.array([r2_dict[d]["r2"] for d in DELTAS])
sc = _one(axes[2], rhos_avg, sigs_j, errs_j, r2s_avg,
xlabel=r"avg auto-correlation $\bar{\rho}$",
ylabel="SIGReg (zscored, 2D joint)",
title=names[2])
plt.colorbar(sc, ax=axes[2], label=r"$R^2$ (avg)")
plt.tight_layout()
plt.savefig(save_path, dpi=300, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# Measurement pipeline
# ═════════════════════════════════════════════════════════════════════════════
def compute_all_transitions(episodes, n_draws=20):
"""
For each delta, compute SIGReg stats and per-dim rho on the
transition-difference distribution z(t+delta) - z(t).
"""
trans = {}
for d in DELTAS:
diffs = (episodes[:, d:] - episodes[:, :-d]).reshape(-1, 2)
rho = np.array([
pearsonr(episodes[:, d:, k].flatten(),
episodes[:, :-d, k].flatten())[0]
for k in range(2)
])
trans[d] = {"rho": rho, "sig": measure(diffs, n_draws=n_draws)}
return trans
# ═════════════════════════════════════════════════════════════════════════════
# Main
# ═════════════════════════════════════════════════════════════════════════════
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--results_dir", type=str, default="results/reacher",
help="Directory with result.json files")
parser.add_argument("--data_path", type=str,
default="data/reacher.h5",
help="HDF5 file with 'qpos' and 'ep_len' datasets "
"(reshaped to (n_episodes, T, 2))")
parser.add_argument("--out_dir", type=str, default="figures/reacher")
parser.add_argument("--n_draws", type=int, default=20,
help="Random subsamples for SIGReg stats")
parser.add_argument("--agg", type=str, default="median",
choices=["median", "mean"],
help="How to aggregate R² across seeds")
args = parser.parse_args()
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
# load episodes (HDF5 with qpos + ep_len, matching notebook convention)
import h5py
with h5py.File(args.data_path, "r") as f:
qpos = np.array(f["qpos"])
ep_len = np.array(f["ep_len"])
T = int(ep_len[0])
episodes = qpos.reshape(-1, T, 2)
print(f"Loaded {len(episodes)} episodes of length {T} from {args.data_path}")
r2_dict = get_best_r2_per_delta(args.results_dir, agg=args.agg)
print(f"Loaded R² for deltas: {sorted(r2_dict.keys())}")
for d in DELTAS:
v = r2_dict[d]
print(f" Δ={d:2d} λ={v['lamb']:.0e} n={v['n_seeds']} "
f"R²={v['r2']:.3f} dim0={v['r2_dim0']:.3f} "
f"dim1={v['r2_dim1']:.3f}")
ou_rho = best_ou_rho(args.results_dir)
if ou_rho is not None:
print(f"Best OU rho: {ou_rho}")
# measure transitions
print("Computing SIGReg on transitions...")
trans = compute_all_transitions(episodes, n_draws=args.n_draws)
# figures
make_scatter_grid(episodes, r2_dict,
save_path=out_dir / "distribution.png")
make_rho_vs_sigreg(trans, r2_dict,
save_path=out_dir / "rho_vs_sigreg.png",
best_ou_rho_val=ou_rho)
if __name__ == "__main__":
main()