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,503 @@
"""
Aggregate and plot Reacher sweep results.
Generates:
1. OU: R² vs ρ (per lambda)
2. OU: R² vs ρ (per dimension)
3. Traj: per-dim R² vs δ (with per-dim ρ annotations)
4. OU vs Traj on same axes (ρ on x-axis, traj uses measured ρ_mean)
5. Lambda robustness panel (OU, R² vs λ for each ρ)
6. Traj: R² vs measured ρ
7. Orthogonality error plots
Usage:
python analysis/plot_reacher.py --results_dir results/reacher
"""
import os
import json
import argparse
import numpy as np
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from collections import defaultdict
# ═════════════════════════════════════════════════════════════════════════════
# DATA LOADING
# ═════════════════════════════════════════════════════════════════════════════
def load_summaries(results_dir):
"""Load all summary_*.json files, split into OU and traj."""
ou_results, traj_results = [], []
for p in sorted(Path(results_dir).glob("summary_*.json")):
with open(p) as f:
summary = json.load(f)
for run_name, r in summary.items():
if r.get("r2_hz", -999) <= -1:
continue
if "delta" in r:
traj_results.append(r)
else:
ou_results.append(r)
return ou_results, traj_results
def _group_by(results, x_key):
"""Group results by (x_key, lamb) → list of result dicts."""
grouped = defaultdict(list)
for r in results:
grouped[(r[x_key], r["lamb"])].append(r)
return grouped
def _get_sorted(results, key):
return sorted(set(r[key] for r in results))
# ═════════════════════════════════════════════════════════════════════════════
# 1. OU: R² vs ρ (one curve per lambda)
# ═════════════════════════════════════════════════════════════════════════════
def plot_ou_r2_vs_rho(results, save_path):
grouped = _group_by(results, "rho")
lambs = _get_sorted(results, "lamb")
rhos = _get_sorted(results, "rho")
fig, ax = plt.subplots(figsize=(7, 4.5))
for lamb in lambs:
medians, q25, q75, xs = [], [], [], []
for rho in rhos:
vals = [r["r2_hz"] for r in grouped.get((rho, lamb), [])]
if vals:
medians.append(np.median(vals))
q25.append(np.percentile(vals, 25))
q75.append(np.percentile(vals, 75))
xs.append(rho)
medians, q25, q75 = np.array(medians), np.array(q25), np.array(q75)
ax.plot(xs, medians, "o-", label=f"λ={lamb:.0e}", markersize=5)
ax.fill_between(xs, q25, q75, alpha=0.15)
ax.set_xlabel("ρ (OU autocorrelation)", fontsize=12)
ax.set_ylabel("R² (embed → true state)", fontsize=12)
ax.set_title("OU: Linear identifiability vs. ρ", fontsize=13)
ax.legend(fontsize=9)
ax.set_ylim(-0.05, 1.05)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# 2. OU: per-dimension R² vs ρ
# ═════════════════════════════════════════════════════════════════════════════
def plot_ou_perdim_r2(results, save_path):
"""Two curves: shoulder vs wrist, averaged over lambda and seeds."""
by_rho = defaultdict(list)
for r in results:
if "r2_hz_per_dim" in r:
by_rho[r["rho"]].append(r["r2_hz_per_dim"])
rhos = sorted(by_rho.keys())
dim0_med, dim1_med = [], []
for rho in rhos:
vals = np.array(by_rho[rho])
dim0_med.append(np.median(vals[:, 0]))
dim1_med.append(np.median(vals[:, 1]))
fig, ax = plt.subplots(figsize=(7, 4.5))
ax.plot(rhos, dim0_med, "o-", label="Shoulder (dim 0)", markersize=5)
ax.plot(rhos, dim1_med, "s-", label="Wrist (dim 1)", markersize=5)
ax.set_xlabel("ρ (OU autocorrelation)", fontsize=12)
ax.set_ylabel("R² per dimension", fontsize=12)
ax.set_title("OU: Per-dimension identifiability", fontsize=13)
ax.legend(fontsize=10)
ax.set_ylim(-0.05, 1.05)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# 3. Traj: per-dimension R² vs δ (with ρ annotations)
# ═════════════════════════════════════════════════════════════════════════════
def plot_traj_perdim_r2(results, save_path):
"""Two curves: shoulder vs wrist, annotated with per-dim ρ."""
by_delta = defaultdict(list)
for r in results:
if "r2_hz_per_dim" in r:
by_delta[r["delta"]].append(r)
deltas = sorted(by_delta.keys())
dim0_med, dim1_med = [], []
rho0_vals, rho1_vals = [], []
for delta in deltas:
runs = by_delta[delta]
vals = np.array([r["r2_hz_per_dim"] for r in runs])
dim0_med.append(np.median(vals[:, 0]))
dim1_med.append(np.median(vals[:, 1]))
rho0 = [r.get("rho_shoulder", None) for r in runs]
rho1 = [r.get("rho_wrist", None) for r in runs]
rho0_vals.append(rho0[0] if rho0[0] is not None else None)
rho1_vals.append(rho1[0] if rho1[0] is not None else None)
fig, ax = plt.subplots(figsize=(8, 5))
ax.plot(deltas, dim0_med, "o-", label="Shoulder (dim 0)",
markersize=5, color="C0")
ax.plot(deltas, dim1_med, "s-", label="Wrist (dim 1)",
markersize=5, color="C1")
# Annotate with ρ values
for i, delta in enumerate(deltas):
if rho0_vals[i] is not None:
y_pos = max(dim0_med[i], dim1_med[i]) + 0.03
ax.annotate(f"ρ₀={rho0_vals[i]:.3f}\nρ₁={rho1_vals[i]:.3f}",
(delta, y_pos),
fontsize=7, ha="center", alpha=0.7)
ax.set_xlabel("δ (temporal stride)", fontsize=12)
ax.set_ylabel("R² per dimension", fontsize=12)
ax.set_title("Trajectory: Per-dimension identifiability vs. δ", fontsize=13)
ax.legend(fontsize=10)
ax.set_ylim(-0.15, 1.15)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# 4. OU vs Traj on same axes (ρ on x-axis)
# ═════════════════════════════════════════════════════════════════════════════
def plot_ou_vs_traj(ou_results, traj_results, save_path):
"""Both conditions on one plot, using ρ as common x-axis."""
# OU: group by rho, average over lambda and seeds
ou_by_rho = defaultdict(list)
for r in ou_results:
ou_by_rho[r["rho"]].append(r["r2_hz"])
# Traj: use rho_mean, group by delta
traj_by_rho = defaultdict(list)
for r in traj_results:
rho_mean = r.get("rho_mean", None)
if rho_mean is not None:
traj_by_rho[rho_mean].append(r["r2_hz"])
fig, ax = plt.subplots(figsize=(7, 4.5))
# OU
rhos_ou = sorted(ou_by_rho.keys())
med_ou = [np.median(ou_by_rho[rho]) for rho in rhos_ou]
q25_ou = [np.percentile(ou_by_rho[rho], 25) for rho in rhos_ou]
q75_ou = [np.percentile(ou_by_rho[rho], 75) for rho in rhos_ou]
ax.plot(rhos_ou, med_ou, "o-", label="OU (Gaussian)", markersize=6,
color="C0", linewidth=2)
ax.fill_between(rhos_ou, q25_ou, q75_ou, alpha=0.15, color="C0")
# Traj
rhos_traj = sorted(traj_by_rho.keys())
med_traj = [np.median(traj_by_rho[rho]) for rho in rhos_traj]
q25_traj = [np.percentile(traj_by_rho[rho], 25) for rho in rhos_traj]
q75_traj = [np.percentile(traj_by_rho[rho], 75) for rho in rhos_traj]
ax.plot(rhos_traj, med_traj, "s-", label="Trajectory (non-Gaussian)",
markersize=6, color="C3", linewidth=2)
ax.fill_between(rhos_traj, q25_traj, q75_traj, alpha=0.15, color="C3")
ax.set_xlabel("ρ (autocorrelation)", fontsize=12)
ax.set_ylabel("R² (embed → true state)", fontsize=12)
ax.set_title("Gaussian vs. non-Gaussian latents", fontsize=13)
ax.legend(fontsize=10)
ax.set_ylim(-0.05, 1.05)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# 5. Lambda robustness (OU)
# ═════════════════════════════════════════════════════════════════════════════
def plot_lambda_robustness(results, save_path):
"""R² vs lambda for each rho."""
grouped = _group_by(results, "rho")
rhos = _get_sorted(results, "rho")
lambs = _get_sorted(results, "lamb")
fig, ax = plt.subplots(figsize=(7, 4.5))
for rho in rhos:
medians, xs = [], []
for lamb in lambs:
vals = [r["r2_hz"] for r in grouped.get((rho, lamb), [])]
if vals:
medians.append(np.median(vals))
xs.append(lamb)
if medians:
ax.plot(xs, medians, "o-", label=f"ρ={rho}", markersize=4)
ax.set_xscale("log")
ax.set_xlabel("λ (SIGReg weight)", fontsize=12)
ax.set_ylabel("R² (embed → true state)", fontsize=12)
ax.set_title("OU: Robustness to λ", fontsize=13)
ax.legend(fontsize=8, ncol=2)
ax.set_ylim(-0.05, 1.05)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# 6. Traj: R² vs measured ρ (per lambda)
# ═════════════════════════════════════════════════════════════════════════════
def plot_traj_r2_vs_rho(results, save_path):
"""Traj R² plotted against measured autocorrelation, per lambda."""
grouped = defaultdict(list)
for r in results:
rho_mean = r.get("rho_mean", None)
if rho_mean is not None:
grouped[(rho_mean, r["lamb"])].append(r["r2_hz"])
lambs = _get_sorted(results, "lamb")
rhos = sorted(set(k[0] for k in grouped.keys()))
fig, ax = plt.subplots(figsize=(7, 4.5))
for lamb in lambs:
medians, xs = [], []
for rho in rhos:
vals = grouped.get((rho, lamb), [])
if vals:
medians.append(np.median(vals))
xs.append(rho)
if medians:
ax.plot(xs, medians, "o-", label=f"λ={lamb:.0e}", markersize=5)
ax.set_xlabel("ρ (measured autocorrelation)", fontsize=12)
ax.set_ylabel("R² (embed → true state)", fontsize=12)
ax.set_title("Trajectory: Identifiability vs. measured ρ", fontsize=13)
ax.legend(fontsize=9)
ax.set_ylim(-0.05, 1.05)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# 7. Orthogonality error
# ═════════════════════════════════════════════════════════════════════════════
def plot_orth_err(results, x_key, x_label, title, save_path):
grouped = _group_by(results, x_key)
lambs = _get_sorted(results, "lamb")
x_vals = _get_sorted(results, x_key)
fig, ax = plt.subplots(figsize=(7, 4.5))
for lamb in lambs:
medians, xs = [], []
for x in x_vals:
vals = [r.get("orth_error", r.get("procrustes_error", 1.0))
for r in grouped.get((x, lamb), [])]
if vals:
medians.append(np.median(vals))
xs.append(x)
ax.plot(xs, medians, "o-", label=f"λ={lamb:.0e}", markersize=4)
ax.set_xlabel(x_label, fontsize=12)
ax.set_ylabel("Orthogonality error", fontsize=12)
ax.set_title(title, fontsize=13)
ax.legend(fontsize=9)
ax.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(save_path, dpi=200, bbox_inches="tight")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# TABLE
# ═════════════════════════════════════════════════════════════════════════════
def print_table(results, x_key, label):
grouped = defaultdict(list)
for r in results:
grouped[(r[x_key], r["lamb"])].append(r)
print(f"\n=== {label} ===")
print(f"{x_key:>6s} {'lamb':>8s} {'R²(h→z)':>10s} "
f"{'R²(dim0)':>10s} {'R²(dim1)':>10s} "
f"{'R²(sincos)':>10s} {'orth_err':>10s} {'n':>4s}")
print("-" * 76)
for (x, lamb), runs in sorted(grouped.items()):
r2s = [r["r2_hz"] for r in runs]
errs = [r.get("orth_error", r.get("procrustes_error", 1.0))
for r in runs]
dim0 = [r["r2_hz_per_dim"][0] for r in runs
if "r2_hz_per_dim" in r]
dim1 = [r["r2_hz_per_dim"][1] for r in runs
if "r2_hz_per_dim" in r]
sincos = [r["r2_sincos"] for r in runs if "r2_sincos" in r]
d0 = f"{np.median(dim0):10.4f}" if dim0 else f"{'n/a':>10s}"
d1 = f"{np.median(dim1):10.4f}" if dim1 else f"{'n/a':>10s}"
sc = f"{np.median(sincos):10.4f}" if sincos else f"{'n/a':>10s}"
print(f"{x:6g} {lamb:8.1e} "
f"{np.median(r2s):10.4f} "
f"{d0} {d1} "
f"{sc} "
f"{np.median(errs):10.4f} "
f"{len(runs):4d}")
def plot_traj_distributions(h5_path, results_dir, save_path):
"""Marginal + transition distributions for trajectory data, annotated with R²."""
import h5py
from scipy.stats import pearsonr
with h5py.File(h5_path, "r") as f:
qpos = np.array(f["qpos"])
ep_len = np.array(f["ep_len"])
T = ep_len[0]
episodes = qpos.reshape(-1, T, 2)
# Load R² per delta
grouped = 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
grouped[(r["delta"], r["lamb"])].append(r)
r2_dict = {}
for delta in set(k[0] for k in grouped):
best_mean, best_lamb = -np.inf, None
for lamb in set(k[1] for k in grouped if k[0] == delta):
m = np.mean([r["r2_hz"] for r in grouped[(delta, lamb)]])
if m > best_mean:
best_mean, best_lamb = m, lamb
runs = grouped[(delta, best_lamb)]
r2_dict[delta] = {
"r2_dim0": np.mean([r["r2_hz_per_dim"][0] for r in runs]),
"r2_dim1": np.mean([r["r2_hz_per_dim"][1] for r in runs]),
}
DELTAS = [1, 2, 4, 8, 16, 32, 64]
s = 0.0001
sub = 1
fig = plt.figure(figsize=0.85 * np.array((1 + 3 * len(DELTAS), 5)))
gs = fig.add_gridspec(2, 2 + len(DELTAS))
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):
# Transition scatter
ax = fig.add_subplot(gs[0, 2 + i])
transitions = (episodes[:, delta:] - episodes[:, :-delta]).reshape(-1, 2)
ax.scatter(*transitions[::sub].T, s=s)
ax.set_title(
r"$\Delta=%d$" % delta + "\n"
+ r"$R^2=(%.2f,\,%.2f)$"
% (r2_dict[delta]["r2_dim0"], r2_dict[delta]["r2_dim1"])
)
ax.grid()
# Autocorrelation scatter
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", format="jpg")
plt.close()
print(f"Saved {save_path}")
# ═════════════════════════════════════════════════════════════════════════════
# MAIN
# ═════════════════════════════════════════════════════════════════════════════
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--results_dir", type=str, default="results/reacher")
parser.add_argument("--out_dir", type=str, default="figures/reacher")
parser.add_argument("--h5_path", type=str, default="data/reacher.h5")
args = parser.parse_args()
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
ou_results, traj_results = load_summaries(args.results_dir)
print(f"Loaded {len(ou_results)} OU runs, {len(traj_results)} traj runs")
# Tables
if ou_results:
print_table(ou_results, "rho", "OU")
if traj_results:
print_table(traj_results, "delta", "Trajectory")
# OU plots
if ou_results:
plot_ou_r2_vs_rho(ou_results, out_dir / "ou_r2_vs_rho.png")
plot_ou_perdim_r2(ou_results, out_dir / "ou_perdim_r2.png")
plot_lambda_robustness(ou_results, out_dir / "ou_lambda_robustness.png")
plot_orth_err(ou_results, "rho", "ρ",
"OU: Orthogonality error vs. ρ",
out_dir / "ou_orth_err.png")
# Traj plots
if traj_results:
plot_traj_perdim_r2(traj_results, out_dir / "traj_perdim_r2.png")
plot_traj_r2_vs_rho(traj_results, out_dir / "traj_r2_vs_rho.png")
plot_orth_err(traj_results, "delta", "δ",
"Trajectory: Orthogonality error vs. δ",
out_dir / "traj_orth_err.png")
h5_path = os.path.join(args.h5_path)
if os.path.exists(h5_path):
plot_traj_distributions(h5_path, args.results_dir,
out_dir / "traj_distributions.png")
# Combined
if ou_results and traj_results:
plot_ou_vs_traj(ou_results, traj_results,
out_dir / "ou_vs_traj.png")
if __name__ == "__main__":
main()