- 移除 JEPA/lejepa-identifiability 子模块 gitlink - 移除 research/multiply/MultiPLY 子模块 gitlink - 删除 .gitmodules(不再有外部 URL 依赖) - 两个目录内容作为普通文件纳入主仓库追踪 - 删除各自内部 .git 目录,消除嵌套 git 仓库
This commit is contained in:
@@ -0,0 +1,63 @@
|
||||
"""
|
||||
Aggregate results from any experiment into a flat CSV.
|
||||
|
||||
Usage:
|
||||
python analysis/aggregate.py --results_dir results/2d/ --out results/2d/summary.csv
|
||||
python analysis/aggregate.py --results_dir results/ --recursive --out results/all.csv
|
||||
"""
|
||||
|
||||
import argparse, glob, json, os
|
||||
import pandas as pd
|
||||
|
||||
|
||||
SCALAR_KEYS = [
|
||||
"experiment", "run_name", "mixing", "encoder", "mode", "source_dist",
|
||||
"seed", "N", "lamb", "rho", "lr", "steps", "batch_size", "n_layers", "hidden",
|
||||
"r2_zx", "r2_xz", "r2_zh", "r2_hz",
|
||||
"orth_err", "orth_err_normalized",
|
||||
"epsilon", "delta", "D_bound", "approx_bound",
|
||||
"procrustes_mse", "L_h", "trace_cov",
|
||||
"final_align", "final_sigreg", "final_whiten", "final_loss",
|
||||
]
|
||||
|
||||
|
||||
def load_results(results_dir, recursive=False):
|
||||
pattern = os.path.join(results_dir, "**/*.json") if recursive else os.path.join(results_dir, "*.json")
|
||||
files = sorted(glob.glob(pattern, recursive=recursive))
|
||||
print(f"Found {len(files)} .json files")
|
||||
|
||||
rows = []
|
||||
for path in files:
|
||||
try:
|
||||
with open(path) as f:
|
||||
r = json.load(f)
|
||||
row = {k: r.get(k) for k in SCALAR_KEYS}
|
||||
row["file"] = os.path.relpath(path, results_dir)
|
||||
rows.append(row)
|
||||
except Exception as e:
|
||||
print(f" SKIP {path}: {e}")
|
||||
return pd.DataFrame(rows)
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dir", type=str, required=True)
|
||||
p.add_argument("--out", type=str, default=None)
|
||||
p.add_argument("--recursive", action="store_true")
|
||||
args = p.parse_args()
|
||||
|
||||
df = load_results(args.results_dir, recursive=args.recursive)
|
||||
if len(df) == 0:
|
||||
print("No results found.")
|
||||
return
|
||||
|
||||
print(f"\n{len(df)} runs loaded")
|
||||
print(df.to_string(index=False))
|
||||
|
||||
out = args.out or os.path.join(args.results_dir, "summary.csv")
|
||||
df.to_csv(out, index=False)
|
||||
print(f"\nSaved {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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 (Epps–Pulley) — 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()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,98 @@
|
||||
"""
|
||||
Ablation table: 4 mixings × 3 source distributions × {SIGReg, Whitening}.
|
||||
Reports R²(h→z). The winner per (mixing, α) is bolded when the difference
|
||||
is statistically significant (Welch's t-test, p < 0.05).
|
||||
|
||||
Usage:
|
||||
python analysis/make_table_ablation.py --out figures/tab_ablation.tex
|
||||
"""
|
||||
import argparse, glob, json, os
|
||||
import numpy as np
|
||||
from scipy import stats
|
||||
from collections import defaultdict
|
||||
|
||||
MIXINGS = [("spiral", "Spiral"), ("banana", "Banana"),
|
||||
("sinusoid", "Sinusoid"), ("nvp", "NVP")]
|
||||
ALPHAS = [
|
||||
(None, r"Gaussian ($\alpha = 2$)"),
|
||||
(0.25, r"Heavy tail / sparse ($\alpha = 1/4$)"),
|
||||
(16.0, r"Light tail / uniform ($\alpha = 16$)"),
|
||||
]
|
||||
# METRIC = "r2_hz" # marginal
|
||||
METRIC = "r2_hz_grid" # grid
|
||||
P_THRESH = 0.05
|
||||
|
||||
|
||||
def alpha_key(r):
|
||||
if r.get("source_dist") == "gennorm":
|
||||
return r.get("source_alpha")
|
||||
return None
|
||||
|
||||
|
||||
def fmt(vals, bold=False):
|
||||
if not vals:
|
||||
return r"$-$"
|
||||
s = f"{np.mean(vals):.3f} \\pm {np.std(vals):.3f}"
|
||||
return rf"$\mathbf{{{s}}}$" if bold else f"${s}$"
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--dirs", nargs="+",
|
||||
default=["results/2d",
|
||||
"results/ablation_alpha_0.25",
|
||||
"results/ablation_alpha_16"])
|
||||
p.add_argument("--out", default="figures/tab_ablation.tex")
|
||||
args = p.parse_args()
|
||||
|
||||
groups = defaultdict(list)
|
||||
for d in args.dirs:
|
||||
for path in sorted(glob.glob(os.path.join(d, "*.json"))):
|
||||
with open(path) as f:
|
||||
r = json.load(f)
|
||||
if METRIC not in r:
|
||||
continue
|
||||
key = (r["mixing"], alpha_key(r), r["mode"])
|
||||
groups[key].append(r[METRIC])
|
||||
|
||||
col_spec = "l" + " cc" * len(ALPHAS)
|
||||
multicol = " & ".join(rf"\multicolumn{{2}}{{c}}{{{lab}}}" for _, lab in ALPHAS)
|
||||
cmidrules = "".join(rf"\cmidrule(lr){{{2*i+2}-{2*i+3}}}" for i in range(len(ALPHAS)))
|
||||
method_hdr = " & ".join(["SIGReg & Whitening"] * len(ALPHAS))
|
||||
|
||||
lines = [
|
||||
rf"\begin{{tabular}}{{{col_spec}}}",
|
||||
r"\toprule",
|
||||
rf" & {multicol} \\",
|
||||
cmidrules,
|
||||
rf"Mixing & {method_hdr} \\",
|
||||
r"\midrule",
|
||||
]
|
||||
for mix_key, mix_name in MIXINGS:
|
||||
row = [mix_name]
|
||||
for alpha_k, _ in ALPHAS:
|
||||
v_lej = groups.get((mix_key, alpha_k, "lejepa"), [])
|
||||
v_wht = groups.get((mix_key, alpha_k, "whiten"), [])
|
||||
bold_lej = bold_wht = False
|
||||
if len(v_lej) >= 2 and len(v_wht) >= 2:
|
||||
_, pval = stats.ttest_ind(v_lej, v_wht, equal_var=False)
|
||||
if pval < P_THRESH:
|
||||
if np.mean(v_lej) > np.mean(v_wht):
|
||||
bold_lej = True
|
||||
else:
|
||||
bold_wht = True
|
||||
row.append(fmt(v_lej, bold=bold_lej))
|
||||
row.append(fmt(v_wht, bold=bold_wht))
|
||||
lines.append(" & ".join(row) + r" \\")
|
||||
lines += [r"\bottomrule", r"\end{tabular}"]
|
||||
|
||||
out = "\n".join(lines)
|
||||
print(out)
|
||||
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
|
||||
with open(args.out, "w") as f:
|
||||
f.write(out + "\n")
|
||||
print(f"\nSaved {args.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,236 @@
|
||||
"""
|
||||
Generate LaTeX tables for the paper, matching scaling table style.
|
||||
|
||||
Usage:
|
||||
python analysis/make_table_reacher.py --results_dir results/reacher
|
||||
"""
|
||||
|
||||
import json
|
||||
import argparse
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
|
||||
|
||||
def load_all_results(results_dir):
|
||||
ou, traj = [], []
|
||||
for p in Path(results_dir).rglob("result.json"):
|
||||
r = json.load(open(p))
|
||||
if "delta" in r and r.get("rho") is None:
|
||||
traj.append(r)
|
||||
elif "rho" in r:
|
||||
ou.append(r)
|
||||
return ou, traj
|
||||
|
||||
|
||||
def best_lambda_per_x(results, x_key):
|
||||
grouped = defaultdict(list)
|
||||
for r in results:
|
||||
grouped[(r[x_key], r["lamb"])].append(r)
|
||||
|
||||
best = {}
|
||||
for x in sorted(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] == x):
|
||||
m = np.mean([r["r2_hz"] for r in grouped[(x, lamb)]])
|
||||
if m > best_mean:
|
||||
best_mean, best_lamb = m, lamb
|
||||
best[x] = {
|
||||
"lamb": best_lamb,
|
||||
"runs": grouped[(x, best_lamb)],
|
||||
}
|
||||
return best
|
||||
|
||||
|
||||
def pm(vals, fmt=".2f"):
|
||||
"""Format as value\\tiny{±std} matching paper style."""
|
||||
m, s = np.mean(vals), np.std(vals)
|
||||
return f"{m:{fmt}}\\tiny{{$\\pm${s:.0e}}}"
|
||||
|
||||
|
||||
def make_combined_table(ou_results, traj_results):
|
||||
ou_best = best_lambda_per_x(ou_results, "rho")
|
||||
traj_best = best_lambda_per_x(traj_results, "delta")
|
||||
|
||||
lines = []
|
||||
lines.append(r"\begin{table}[t]")
|
||||
lines.append(r"\centering")
|
||||
lines.append(r"\caption{")
|
||||
lines.append(r" \textbf{Pixel-observation identifiability on DMC Reacher} "
|
||||
r"(mean $\pm$ std over 3 seeds, best $\lambda$ per condition).")
|
||||
lines.append(r" \textbf{Left:} OU process with Gaussian marginals. "
|
||||
r"$R^2$ increases monotonically with $\rho$, reaching $0.95$ "
|
||||
r"at $\rho = 0.99$, confirming linear identifiability from pixels.")
|
||||
lines.append(r" \textbf{Right:} Real SAC trajectories with non-Gaussian marginals. "
|
||||
r"The two joints have different autocorrelation timescales ($\rho_0 \neq \rho_1$) "
|
||||
r"and the wrist has a near-uniform marginal distribution, "
|
||||
r"leading to anisotropic and reduced identifiability.")
|
||||
lines.append(r"}")
|
||||
lines.append(r"\label{tab:reacher}")
|
||||
lines.append(r"\resizebox{\textwidth}{!}{%")
|
||||
lines.append(r"\begin{tabular}{r cc | r cc ccc}")
|
||||
lines.append(r" \multicolumn{3}{c}{\textbf{OU (Gaussian)}} & "
|
||||
r"\multicolumn{6}{c}{\textbf{Trajectory (non-Gaussian)}} \\")
|
||||
lines.append(r"\cmidrule(lr){1-3} \cmidrule(lr){4-9}")
|
||||
lines.append(r"$\rho$ & $R^2(z \to h)$ & $R^2(h \to z)$ & "
|
||||
r"$\delta$ & $\rho_0$ & $\rho_1$ & "
|
||||
r"$R^2(z \to h)$ & $R^2(h \to z_0)$ & $R^2(h \to z_1)$ \\")
|
||||
lines.append(r"\midrule")
|
||||
|
||||
ou_rhos = sorted(ou_best.keys())
|
||||
traj_deltas = sorted(traj_best.keys())
|
||||
n_rows = max(len(ou_rhos), len(traj_deltas))
|
||||
|
||||
for i in range(n_rows):
|
||||
# OU columns
|
||||
if i < len(ou_rhos):
|
||||
rho = ou_rhos[i]
|
||||
runs = ou_best[rho]["runs"]
|
||||
r2_zh = pm([r["r2_zh"] for r in runs])
|
||||
r2_hz = pm([r["r2_hz"] for r in runs])
|
||||
ou_str = f" {rho:.2f} & {r2_zh} & {r2_hz}"
|
||||
else:
|
||||
ou_str = r" & &"
|
||||
|
||||
# Traj columns
|
||||
if i < len(traj_deltas):
|
||||
delta = traj_deltas[i]
|
||||
runs = traj_best[delta]["runs"]
|
||||
rho0 = runs[0].get("rho_shoulder", None)
|
||||
rho1 = runs[0].get("rho_wrist", None)
|
||||
rho0_s = f"{rho0:.3f}" if rho0 is not None else "---"
|
||||
rho1_s = f"{rho1:.3f}" if rho1 is not None else "---"
|
||||
r2_zh = pm([r["r2_zh"] for r in runs])
|
||||
r2_d0 = pm([r["r2_hz_per_dim"][0] for r in runs])
|
||||
r2_d1 = pm([r["r2_hz_per_dim"][1] for r in runs])
|
||||
traj_str = f"{delta} & {rho0_s} & {rho1_s} & {r2_zh} & {r2_d0} & {r2_d1}"
|
||||
else:
|
||||
traj_str = r"& & & & &"
|
||||
|
||||
lines.append(f"{ou_str} & {traj_str} \\\\")
|
||||
|
||||
lines.append(r"\bottomrule")
|
||||
lines.append(r"\end{tabular}}")
|
||||
lines.append(r"\vspace{5pt}")
|
||||
lines.append(r"\vspace{-20pt}")
|
||||
lines.append(r"\end{table}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def make_ou_table_standalone(ou_results):
|
||||
"""Standalone OU table for appendix if needed."""
|
||||
ou_best = best_lambda_per_x(ou_results, "rho")
|
||||
|
||||
lines = []
|
||||
lines.append(r"\begin{table}[t]")
|
||||
lines.append(r"\centering")
|
||||
lines.append(r"\begin{tabular}{r c cc}")
|
||||
lines.append(r"\toprule")
|
||||
lines.append(r" \multicolumn{1}{c}{\textbf{Correlation}} & "
|
||||
r"\multicolumn{1}{c}{\textbf{Regularizer}} & "
|
||||
r"\multicolumn{2}{c}{\textbf{Linear identifiability}} \\")
|
||||
lines.append(r"\cmidrule(lr){1-1} \cmidrule(lr){2-2} \cmidrule(lr){3-4}")
|
||||
lines.append(r"$\rho$ & $\lambda$ & $R^2(z \to h)$ & $R^2(h \to z)$ \\")
|
||||
lines.append(r"\midrule")
|
||||
|
||||
for rho in sorted(ou_best.keys()):
|
||||
runs = ou_best[rho]["runs"]
|
||||
lamb = ou_best[rho]["lamb"]
|
||||
r2_zh = pm([r["r2_zh"] for r in runs])
|
||||
r2_hz = pm([r["r2_hz"] for r in runs])
|
||||
lines.append(f" {rho:.2f} & {lamb:.0e} & {r2_zh} & {r2_hz} \\\\")
|
||||
|
||||
lines.append(r"\bottomrule")
|
||||
lines.append(r"\end{tabular}")
|
||||
lines.append(r"\vspace{5pt}")
|
||||
lines.append(r"\caption{")
|
||||
lines.append(r" \textbf{OU (Gaussian) identifiability from pixels} "
|
||||
r"(mean $\pm$ std over 3 seeds).")
|
||||
lines.append(r" $R^2$ increases monotonically with temporal correlation $\rho$, "
|
||||
r"reaching $0.95$ at $\rho = 0.99$.")
|
||||
lines.append(r"}")
|
||||
lines.append(r"\label{tab:reacher_ou}")
|
||||
lines.append(r"\end{table}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def make_traj_table_standalone(traj_results):
|
||||
"""Standalone traj table for appendix if needed."""
|
||||
traj_best = best_lambda_per_x(traj_results, "delta")
|
||||
|
||||
lines = []
|
||||
lines.append(r"\begin{table}[t]")
|
||||
lines.append(r"\centering")
|
||||
lines.append(r"\resizebox{\textwidth}{!}{%")
|
||||
lines.append(r"\begin{tabular}{r cc c c cc}")
|
||||
lines.append(r"\toprule")
|
||||
lines.append(r" \multicolumn{1}{c}{\textbf{Stride}} & "
|
||||
r"\multicolumn{2}{c}{\textbf{Autocorrelation}} & "
|
||||
r"\multicolumn{1}{c}{\textbf{Regularizer}} & "
|
||||
r"\multicolumn{1}{c}{\textbf{Identifiability}} & "
|
||||
r"\multicolumn{2}{c}{\textbf{Per-dimension}} \\")
|
||||
lines.append(r"\cmidrule(lr){1-1} \cmidrule(lr){2-3} \cmidrule(lr){4-4} "
|
||||
r"\cmidrule(lr){5-5} \cmidrule(lr){6-7}")
|
||||
lines.append(r"$\delta$ & $\rho_0$ & $\rho_1$ & $\lambda$ & "
|
||||
r"$R^2(z \to h)$ & $R^2(h \to z_0)$ & $R^2(h \to z_1)$ \\")
|
||||
lines.append(r"\midrule")
|
||||
|
||||
for delta in sorted(traj_best.keys()):
|
||||
runs = traj_best[delta]["runs"]
|
||||
lamb = traj_best[delta]["lamb"]
|
||||
rho0 = runs[0].get("rho_shoulder", None)
|
||||
rho1 = runs[0].get("rho_wrist", None)
|
||||
rho0_s = f"{rho0:.3f}" if rho0 is not None else "---"
|
||||
rho1_s = f"{rho1:.3f}" if rho1 is not None else "---"
|
||||
r2_zh = pm([r["r2_zh"] for r in runs])
|
||||
r2_d0 = pm([r["r2_hz_per_dim"][0] for r in runs])
|
||||
r2_d1 = pm([r["r2_hz_per_dim"][1] for r in runs])
|
||||
lines.append(f" {delta} & {rho0_s} & {rho1_s} & {lamb:.0e} "
|
||||
f"& {r2_zh} & {r2_d0} & {r2_d1} \\\\")
|
||||
|
||||
lines.append(r"\bottomrule")
|
||||
lines.append(r"\end{tabular}}")
|
||||
lines.append(r"\vspace{5pt}")
|
||||
lines.append(r"\caption{")
|
||||
lines.append(r" \textbf{Trajectory (non-Gaussian) identifiability from pixels} "
|
||||
r"(mean $\pm$ std over 3 seeds).")
|
||||
lines.append(r" The shoulder ($z_0$) and wrist ($z_1$) have different "
|
||||
r"autocorrelation timescales and marginal distributions, "
|
||||
r"leading to anisotropic identifiability.")
|
||||
lines.append(r"}")
|
||||
lines.append(r"\label{tab:reacher_traj}")
|
||||
lines.append(r"\end{table}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--results_dir", type=str, default="results/reacher")
|
||||
args = parser.parse_args()
|
||||
|
||||
ou, traj = load_all_results(args.results_dir)
|
||||
print(f"Loaded {len(ou)} OU runs, {len(traj)} traj runs\n")
|
||||
|
||||
if ou and traj:
|
||||
print("=" * 70)
|
||||
print("COMBINED TABLE (for main text)")
|
||||
print("=" * 70)
|
||||
print(make_combined_table(ou, traj))
|
||||
print()
|
||||
|
||||
if ou:
|
||||
print("=" * 70)
|
||||
print("OU TABLE (standalone, for appendix)")
|
||||
print("=" * 70)
|
||||
print(make_ou_table_standalone(ou))
|
||||
print()
|
||||
|
||||
if traj:
|
||||
print("=" * 70)
|
||||
print("TRAJ TABLE (standalone, for appendix)")
|
||||
print("=" * 70)
|
||||
print(make_traj_table_standalone(traj))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,272 @@
|
||||
"""
|
||||
Generate LaTeX tables from scaling results.
|
||||
|
||||
Emits four tables:
|
||||
1. tab:scaling-comparison (main text)
|
||||
Three-way R^2(h -> z) comparison: SIGReg, VICReg, InfoNCE.
|
||||
|
||||
2. tab:scaling-sigreg (appendix)
|
||||
Detailed per-method table for SIGReg: mixing difficulty,
|
||||
linear identifiability (both directions), orthogonality error,
|
||||
alignment loss, SIGReg loss.
|
||||
|
||||
3. tab:scaling-vicreg (appendix)
|
||||
Same structure as SIGReg, but with whitening loss column.
|
||||
|
||||
4. tab:scaling-infonce (appendix)
|
||||
Same structure as SIGReg, but with InfoNCE loss column.
|
||||
|
||||
The three appendix tables let each method tell its own failure-mode story:
|
||||
- SIGReg / VICReg: orthogonality error grows gradually with N
|
||||
- InfoNCE: regularizer loss explodes / fails to converge at high N
|
||||
|
||||
Usage:
|
||||
python analysis/make_table_scaling.py --results_dir results/scaling/
|
||||
"""
|
||||
|
||||
import argparse, glob, json, os
|
||||
import numpy as np
|
||||
from collections import defaultdict
|
||||
import math
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Number formatting helpers (shared)
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def column_scale(stds, threshold=0.05):
|
||||
valid = [s for s in stds if s > 0 and not np.isnan(s)]
|
||||
if not valid:
|
||||
return 0
|
||||
m = max(valid)
|
||||
if m >= threshold:
|
||||
return 0
|
||||
return int(np.floor(np.log10(m))) - 1
|
||||
|
||||
|
||||
def fmt_std(s, k):
|
||||
if np.isnan(s):
|
||||
return "---"
|
||||
scaled = s if k == 0 else s / (10 ** k)
|
||||
if scaled >= 10:
|
||||
return f"{math.floor(scaled):.0f}"
|
||||
elif scaled >= 1:
|
||||
return f"{math.floor(scaled * 10) / 10:.1f}"
|
||||
else:
|
||||
return f"{math.floor(scaled * 100) / 100:.2f}"
|
||||
|
||||
|
||||
def scale_header(k):
|
||||
if k == 0:
|
||||
return r"{\scriptsize $\pm$std}"
|
||||
return rf"{{\scriptsize $\pm$std\,$\times 10^{{{k}}}$}}"
|
||||
|
||||
|
||||
def fmt_cell(m, s, dec, k):
|
||||
if np.isnan(m):
|
||||
return "---"
|
||||
factor = 10 ** dec
|
||||
m_floored = math.floor(m * factor) / factor
|
||||
return rf"{m_floored:.{dec}f}\tiny{{$\pm${fmt_std(s, k)}}}"
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Main-text: three-way comparison on R^2(h -> z)
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def render_table_comparison(by_key, dims):
|
||||
modes = ("lejepa", "whiten", "infonce")
|
||||
|
||||
def agg(N, mode, key):
|
||||
rs = by_key.get((N, mode), [])
|
||||
vals = [r[key] for r in rs if r.get(key) is not None]
|
||||
if not vals:
|
||||
return float("nan"), float("nan")
|
||||
return float(np.mean(vals)), float(np.std(vals))
|
||||
|
||||
def agg_mixing(N, key):
|
||||
all_rs = sum((by_key.get((N, m), []) for m in modes), [])
|
||||
vals = [r[key] for r in all_rs if r.get(key) is not None]
|
||||
if not vals:
|
||||
return float("nan"), float("nan")
|
||||
return float(np.mean(vals)), float(np.std(vals))
|
||||
|
||||
mix_scale = column_scale([agg_mixing(N, "r2_xz")[1] for N in dims])
|
||||
r2_scales = {m: column_scale([agg(N, m, "r2_hz")[1] for N in dims]) for m in modes}
|
||||
|
||||
print(r"\begin{table}[t]")
|
||||
print(r"\centering")
|
||||
print(r"\caption{\textbf{Scaling Comparison Across Regularizers} (mean $\pm$ std, 5 seeds). "
|
||||
r"All three Gaussianity-enforcing objectives are tested on the same RealNVP mixing "
|
||||
r"with matched encoder. SIGReg and VICReg (batch-statistic estimators) maintain "
|
||||
r"$R^2 > 0.999$ up to $N{=}1024$, consistent with Thm.~\ref{thm:approx}. "
|
||||
r"InfoNCE (pair-based) matches at low $N$ but degrades at scale under fixed kernel "
|
||||
r"width $\sigma{=}1$, illustrating the per-dimension tuning required by pair-based estimators. "
|
||||
r"Per-method details (orthogonality, regularizer loss) in App.~\ref{app:scaling}, "
|
||||
r"Tabs.~\ref{tab:scaling-sigreg}--\ref{tab:scaling-infonce}.}")
|
||||
print(r"\label{tab:scaling-comparison}")
|
||||
print(r"\begin{tabular}{r c ccc}")
|
||||
print(r"\toprule")
|
||||
print(r" & \textbf{Mixing} & \multicolumn{3}{c}{\textbf{Linear identifiability} $R^2(h \to z)$} \\")
|
||||
print(r"\cmidrule(lr){3-5}")
|
||||
print(r"$N$ & $R^2(x \to z)$ & SIGReg & VICReg & InfoNCE \\")
|
||||
sub_cells = [
|
||||
scale_header(mix_scale),
|
||||
scale_header(r2_scales["lejepa"]),
|
||||
scale_header(r2_scales["whiten"]),
|
||||
scale_header(r2_scales["infonce"]),
|
||||
]
|
||||
print(" & " + " & ".join(sub_cells) + r" \\")
|
||||
print(r"\midrule")
|
||||
|
||||
for N in dims:
|
||||
cells = [rf"{N}"]
|
||||
m, s = agg_mixing(N, "r2_xz")
|
||||
cells.append(fmt_cell(m, s, 3, mix_scale))
|
||||
for mode in modes:
|
||||
m_, s_ = agg(N, mode, "r2_hz")
|
||||
cells.append(fmt_cell(m_, s_, 6, r2_scales[mode]))
|
||||
print(" " + " & ".join(cells) + r" \\")
|
||||
|
||||
print(r"\bottomrule")
|
||||
print(r"\end{tabular}")
|
||||
print(r"\end{table}")
|
||||
print()
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Appendix: detailed per-method table
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
|
||||
# Each method has its own native regularizer loss key.
|
||||
METHOD_SPECS = {
|
||||
"lejepa": {
|
||||
"name": "SIGReg",
|
||||
"label": "tab:scaling-sigreg",
|
||||
"reg_loss_key": "final_sigreg",
|
||||
"reg_loss_label": "SIGReg",
|
||||
"reg_loss_dec": 2,
|
||||
"caption_tail": (
|
||||
"The RealNVP mixing is consistently nonlinear across dimensions "
|
||||
r"($R^2(x \to z) < 1$). The learned model nonetheless recovers the "
|
||||
"true latents at all dimensions. Training losses are stable; "
|
||||
r"orthogonality error grows gradually with $N$."
|
||||
),
|
||||
},
|
||||
"whiten": {
|
||||
"name": "VICReg",
|
||||
"label": "tab:scaling-vicreg",
|
||||
"reg_loss_key": "final_whiten",
|
||||
"reg_loss_label": "Whitening",
|
||||
"reg_loss_dec": 4,
|
||||
"caption_tail": (
|
||||
"VICReg matches SIGReg on linear identifiability across all dimensions; "
|
||||
r"orthogonality error grows similarly with $N$. The whitening loss is "
|
||||
"stable and small throughout."
|
||||
),
|
||||
},
|
||||
"infonce": {
|
||||
"name": "InfoNCE",
|
||||
"label": "tab:scaling-infonce",
|
||||
"reg_loss_key": "final_loss",
|
||||
"reg_loss_label": "InfoNCE",
|
||||
"reg_loss_dec": 3,
|
||||
"caption_tail": (
|
||||
r"InfoNCE matches the batch-statistic methods at low $N$ but degrades "
|
||||
"at scale under a fixed Gaussian kernel width. The InfoNCE column "
|
||||
"shows the total contrastive loss (not decomposable into alignment "
|
||||
"plus regularizer), which inflates with $N$ as the per-dimension "
|
||||
"kernel-width assumption breaks down."
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def render_table_per_method(by_key, dims, mode):
|
||||
spec = METHOD_SPECS[mode]
|
||||
rows = sum((by_key.get((N, mode), []) for N in dims), [])
|
||||
if not rows:
|
||||
print(f"% No rows for mode={mode}, skipping {spec['label']}")
|
||||
return
|
||||
|
||||
cols = [
|
||||
("r2_zx", 3, r"$R^2(z \to x)$"),
|
||||
("r2_xz", 3, r"$R^2(x \to z)$"),
|
||||
("r2_zh", 5, r"$R^2(z \to h)$"),
|
||||
("r2_hz", 5, r"$R^2(h \to z)$"),
|
||||
("orth_err_normalized", 3, r"$\|\hat Q^\top \hat Q - I\|_F / \sqrt{N}$"),
|
||||
("final_align", 4, "Align"),
|
||||
(spec["reg_loss_key"], spec["reg_loss_dec"], spec["reg_loss_label"]),
|
||||
]
|
||||
|
||||
agg = {}
|
||||
for N in dims:
|
||||
rs = [r for r in rows if r["N"] == N]
|
||||
for key, _, _ in cols:
|
||||
vals = [r[key] for r in rs if r.get(key) is not None]
|
||||
agg[(N, key)] = ((float(np.mean(vals)), float(np.std(vals)))
|
||||
if vals else (float("nan"), float("nan")))
|
||||
|
||||
scales = {key: column_scale([agg[(N, key)][1] for N in dims])
|
||||
for key, _, _ in cols}
|
||||
|
||||
print(r"\begin{table}[t]")
|
||||
print(r"\centering")
|
||||
print(rf"\caption{{\textbf{{Scaling Experiment ({spec['name']})}} "
|
||||
r"(mean $\pm$ std, 5 seeds). " + spec["caption_tail"] + "}")
|
||||
print(rf"\label{{{spec['label']}}}")
|
||||
print(r"\resizebox{\textwidth}{!}{%")
|
||||
print(r"\begin{tabular}{r cc cc c cc}")
|
||||
print(r"\toprule")
|
||||
print(r" \multicolumn{1}{c}{\textbf{Latents}} "
|
||||
r"& \multicolumn{2}{c}{\textbf{Mixing difficulty}} "
|
||||
r"& \multicolumn{2}{c}{\textbf{Linear identifiability}} "
|
||||
r"& \multicolumn{1}{c}{\textbf{Orthogonality}} "
|
||||
rf"& \multicolumn{{2}}{{c}}{{\textbf{{{spec['name']} losses}}}} \\")
|
||||
print(r"\cmidrule(lr){1-1} \cmidrule(lr){2-3} \cmidrule(lr){4-5} "
|
||||
r"\cmidrule(lr){6-6} \cmidrule(lr){7-8}")
|
||||
print("$N$ & " + " & ".join(label for _, _, label in cols) + r" \\")
|
||||
print(" & " + " & ".join(scale_header(scales[key]) for key, _, _ in cols) + r" \\")
|
||||
print(r"\midrule")
|
||||
|
||||
for N in dims:
|
||||
log2N = int(np.log2(N))
|
||||
cells = [rf"$2^{{{log2N}}}$"]
|
||||
for key, dec, _ in cols:
|
||||
m, s = agg[(N, key)]
|
||||
cells.append(fmt_cell(m, s, dec, scales[key]))
|
||||
print(" " + " & ".join(cells) + r" \\")
|
||||
|
||||
print(r"\bottomrule")
|
||||
print(r"\end{tabular}}")
|
||||
print(r"\end{table}")
|
||||
print()
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Driver
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dir", default="results/scaling/")
|
||||
args = p.parse_args()
|
||||
|
||||
by_key = defaultdict(list)
|
||||
for path in sorted(glob.glob(os.path.join(args.results_dir, "*.json"))):
|
||||
with open(path) as f:
|
||||
r = json.load(f)
|
||||
mode = r.get("mode", "lejepa")
|
||||
by_key[(r["N"], mode)].append(r)
|
||||
|
||||
dims = sorted({N for (N, _) in by_key})
|
||||
modes_present = sorted({mode for (_, mode) in by_key})
|
||||
print(f"% Found modes: {modes_present}, dims: {dims}\n")
|
||||
|
||||
print(r"% ── Main text: three-way comparison ──")
|
||||
render_table_comparison(by_key, dims)
|
||||
|
||||
print(r"% ── Appendix: detailed per-method tables ──")
|
||||
for mode in ("lejepa", "whiten", "infonce"):
|
||||
render_table_per_method(by_key, dims, mode)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,82 @@
|
||||
"""
|
||||
2D illustration figure: one row per mixing (z, g(z), h(z)).
|
||||
Picks best seed per mixing by final_loss.
|
||||
|
||||
Usage:
|
||||
python analysis/plot_2d.py --results_dir results/2d/ --out figures/
|
||||
"""
|
||||
|
||||
import argparse, os, glob
|
||||
import torch
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import colorsys
|
||||
|
||||
MIXING_ORDER = ["spiral", "banana", "sinusoid", "nvp"]
|
||||
|
||||
|
||||
def make_colors(z):
|
||||
x, y = z[:, 0], z[:, 1]
|
||||
angles = np.arctan2(y, x)
|
||||
radii = np.sqrt(x**2 + y**2)
|
||||
hue = (angles + np.pi) / (2 * np.pi)
|
||||
lightness = 0.3 + 0.4 * (radii / (radii.max() + 1e-8))
|
||||
saturation = np.full_like(hue, 0.85)
|
||||
return [colorsys.hls_to_rgb(h, l, s) for h, l, s in zip(hue, lightness, saturation)]
|
||||
|
||||
|
||||
def load_best_per_mixing(results_dir):
|
||||
"""Load all results, pick best lejepa run per mixing by final_loss."""
|
||||
files = sorted(glob.glob(os.path.join(results_dir, "*.pt")))
|
||||
by_mix = {}
|
||||
for path in files:
|
||||
r = torch.load(path, map_location="cpu", weights_only=False)
|
||||
if r.get("mode") != "lejepa":
|
||||
continue
|
||||
mix = r["mixing"]
|
||||
if mix not in by_mix or r["final_loss"] < by_mix[mix]["final_loss"]:
|
||||
by_mix[mix] = r
|
||||
return by_mix
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dir", default="results/2d/")
|
||||
p.add_argument("--out", default="figures/")
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
best = load_best_per_mixing(args.results_dir)
|
||||
|
||||
s, lim = 5, 4
|
||||
for mix_name in MIXING_ORDER:
|
||||
if mix_name not in best:
|
||||
print(f"Missing {mix_name}"); continue
|
||||
res = best[mix_name]
|
||||
z, x, h = res["z"], res["x"], res["h"]
|
||||
colors = make_colors(z)
|
||||
|
||||
fig, axes = plt.subplots(1, 3, figsize=(9, 3))
|
||||
for i, (ax, data, labels) in enumerate(zip(
|
||||
axes, [z, x, h],
|
||||
[("True Latent 0", "True Latent 1"),
|
||||
("Observation 0", "Observation 1"),
|
||||
("Learned Latent 0", "Learned Latent 1")],
|
||||
)):
|
||||
ax.scatter(data[:, 0], data[:, 1], c=colors, s=s, linewidths=0)
|
||||
ax.set_xlabel(labels[0])
|
||||
ax.set_ylabel(labels[1])
|
||||
ax.grid(alpha=0.3)
|
||||
if i == 0 or i == 2:
|
||||
ax.set_xlim(-4, 4)
|
||||
ax.set_ylim(-4, 4)
|
||||
|
||||
fig.tight_layout()
|
||||
out_path = os.path.join(args.out, f"fig_2d_{mix_name}.jpg")
|
||||
fig.savefig(out_path, bbox_inches="tight", dpi=500)
|
||||
print(f"Saved {out_path}")
|
||||
plt.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
Laplace ablation: 4 rows (mixings) x 4 cols (z, g(z), h_lejepa, h_whiten).
|
||||
Picks best seed per (mixing, mode) by final_loss.
|
||||
|
||||
Usage:
|
||||
python analysis/plot_ablation.py --results_dir results/ablation_alpha_0.25/ --prefix ablation_alpha_0.25 --out figures/
|
||||
python analysis/plot_ablation.py --results_dir results/ablation_alpha_16/ --prefix ablation_alpha_16 --out figures/
|
||||
"""
|
||||
|
||||
import argparse, os, glob
|
||||
import torch
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import colorsys
|
||||
|
||||
MIXING_ORDER = ["spiral", "banana", "sinusoid", "nvp"]
|
||||
|
||||
|
||||
def make_colors(z):
|
||||
x, y = z[:, 0], z[:, 1]
|
||||
angles = np.arctan2(y, x)
|
||||
radii = np.sqrt(x**2 + y**2)
|
||||
hue = (angles + np.pi) / (2 * np.pi)
|
||||
lightness = 0.3 + 0.4 * (radii / (radii.max() + 1e-8))
|
||||
saturation = np.full_like(hue, 0.85)
|
||||
return [colorsys.hls_to_rgb(h, l, s) for h, l, s in zip(hue, lightness, saturation)]
|
||||
|
||||
|
||||
def load_best(results_dir):
|
||||
files = sorted(glob.glob(os.path.join(results_dir, "*.pt")))
|
||||
by_key = {}
|
||||
for path in files:
|
||||
r = torch.load(path, map_location="cpu", weights_only=False)
|
||||
key = r["run_name"] # e.g. "spiral_lejepa"
|
||||
if key not in by_key or r["final_loss"] < by_key[key]["final_loss"]:
|
||||
by_key[key] = r
|
||||
return by_key
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dir", default="results/ablation/")
|
||||
p.add_argument("--out", default="figures/")
|
||||
p.add_argument("--prefix", default="ablation",
|
||||
help="Output filename prefix (e.g. ablation_alpha_0.25)")
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
best = load_best(args.results_dir)
|
||||
|
||||
s, lim = 5, 4
|
||||
for mix in MIXING_ORDER:
|
||||
lej_key = f"{mix}_lejepa"
|
||||
wht_key = f"{mix}_whiten"
|
||||
if lej_key not in best or wht_key not in best:
|
||||
print(f"Missing {mix}"); continue
|
||||
|
||||
lej, wht = best[lej_key], best[wht_key]
|
||||
z = lej["z"]
|
||||
colors = make_colors(z)
|
||||
|
||||
fig, axes = plt.subplots(1, 4, figsize=(12, 3))
|
||||
col_labels = [
|
||||
("True Latent 0", "True Latent 1"),
|
||||
("Observation 0", "Observation 1"),
|
||||
("Learned (LeJEPA) 0", "Learned (LeJEPA) 1"),
|
||||
("Learned (Whiten) 0", "Learned (Whiten) 1"),
|
||||
]
|
||||
panels = [z, lej["x"], lej["h"], wht["h"]]
|
||||
# r2s = [None, None, lej["r2_hz"], wht["r2_hz"]]
|
||||
r2s = [None, None, lej["r2_hz_grid"], wht["r2_hz_grid"]]
|
||||
|
||||
for i, (ax, data, labels, r2) in enumerate(zip(axes, panels, col_labels, r2s)):
|
||||
ax.scatter(data[:, 0], data[:, 1], c=colors, s=s, linewidths=0)
|
||||
ax.set_xlabel(labels[0])
|
||||
ax.set_ylabel(labels[1])
|
||||
ax.grid(alpha=0.3)
|
||||
if r2 is not None:
|
||||
ax.text(0.95, 0.05, f"$R^2$={r2:.3f}", transform=ax.transAxes,
|
||||
ha="right", va="bottom", fontsize=9,
|
||||
bbox=dict(boxstyle="round,pad=0.2", fc="white", alpha=0.8))
|
||||
|
||||
fig.tight_layout()
|
||||
out_path = os.path.join(args.out, f"fig_{args.prefix}_{mix}.jpg")
|
||||
fig.savefig(out_path, bbox_inches="tight", dpi=500)
|
||||
print(f"Saved {out_path}")
|
||||
plt.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,237 @@
|
||||
"""
|
||||
Bound verification and grid search plots.
|
||||
|
||||
bound_verification.pdf — pooled across experiments (main paper)
|
||||
bound_decomposition.pdf — grid search only (appendix)
|
||||
heatmap_*.pdf — grid search only (appendix)
|
||||
|
||||
Usage:
|
||||
python analysis/plot_bound.py \
|
||||
--results_dirs results/grid results/2d results/scaling results/gennorm \
|
||||
--out figures/
|
||||
"""
|
||||
|
||||
import argparse, os, glob, json
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib as mpl
|
||||
|
||||
# mpl.rcParams.update({
|
||||
# "font.size": 10, "axes.titlesize": 11, "axes.labelsize": 10,
|
||||
# "figure.dpi": 200, "font.family": "serif",
|
||||
# })
|
||||
|
||||
EXPERIMENT_COLORS = {
|
||||
"2d": "tab:blue",
|
||||
"grid": "tab:red",
|
||||
"scaling": "tab:green",
|
||||
"reacher": "tab:purple",
|
||||
"gennorm": "tab:orange",
|
||||
}
|
||||
EXPERIMENT_MARKERS = {
|
||||
"2d": "o",
|
||||
"grid": "D",
|
||||
"scaling": "s",
|
||||
"reacher": "v",
|
||||
"gennorm": "^",
|
||||
}
|
||||
EXPERIMENT_ORDER = ["grid", "scaling", "2d", "reacher", "gennorm"]
|
||||
|
||||
|
||||
def is_valid_run(r, path=""):
|
||||
"""SIGReg + Gaussian source + non-degenerate (encoder actually learned)."""
|
||||
if r.get("mode") != "lejepa":
|
||||
return False
|
||||
sd = r.get("source_dist", "gaussian")
|
||||
if sd == "gennorm" and abs(r.get("source_alpha", 0) - 2.0) > 1e-6:
|
||||
return False
|
||||
if sd not in ("gaussian", "gennorm"):
|
||||
return False
|
||||
# Drop degenerate runs where the encoder failed to learn
|
||||
if r.get("r2_hz", 0) < 0.5:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def load_all(dirs):
|
||||
data = []
|
||||
for d in dirs:
|
||||
for path in sorted(glob.glob(os.path.join(d, "**", "*.json"), recursive=True)):
|
||||
with open(path) as f:
|
||||
r = json.load(f)
|
||||
if not isinstance(r, dict):
|
||||
continue
|
||||
if r.get("approx_bound") is None or r.get("procrustes_mse") is None:
|
||||
continue
|
||||
if not is_valid_run(r, path):
|
||||
continue
|
||||
data.append(r)
|
||||
print(f"Loaded {len(data)} Gaussian-source runs")
|
||||
return data
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dirs", nargs="+", required=True)
|
||||
p.add_argument("--out", default="figures/")
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
data = load_all(args.results_dirs)
|
||||
if not data:
|
||||
print("No results."); return
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════
|
||||
# Bound verification (main paper, single panel)
|
||||
# ══════════════════════════════════════════════════════════════════════
|
||||
errors = np.array([d["procrustes_mse"] for d in data])
|
||||
bounds = np.array([d["approx_bound"] for d in data])
|
||||
experiments = [d["experiment"] for d in data]
|
||||
|
||||
fig, ax = plt.subplots(figsize=0.8 * np.array((3, 3)))
|
||||
|
||||
pos = (errors > 0) & (bounds > 0)
|
||||
lo = min(errors[pos].min(), bounds[pos].min()) * 0.3
|
||||
hi = max(errors.max(), bounds.max()) * 3
|
||||
pts = np.logspace(np.log10(lo), np.log10(hi), 200)
|
||||
ax.fill_between(pts, lo * 0.1, pts, color='#c8e6c9', alpha=0.35, zorder=0)
|
||||
ax.fill_between(pts, pts, hi * 10, color='#ffcdd2', alpha=0.35, zorder=0)
|
||||
|
||||
for exp in EXPERIMENT_ORDER:
|
||||
mask = np.array([e == exp for e in experiments])
|
||||
if not mask.any():
|
||||
continue
|
||||
ax.scatter(bounds[mask], errors[mask],
|
||||
c=EXPERIMENT_COLORS[exp],
|
||||
marker=EXPERIMENT_MARKERS[exp],
|
||||
s=32, alpha=0.7, edgecolors='k', linewidths=0.3,
|
||||
label=exp, zorder=3)
|
||||
|
||||
ax.plot([lo, hi], [lo, hi], 'k--', alpha=0.5, linewidth=0.8)
|
||||
ax.set_xscale('log'); ax.set_yscale('log')
|
||||
ax.set_xlim(lo, hi)
|
||||
ax.set_ylim(lo, hi)
|
||||
ax.set_xlabel("Recovery error bound")
|
||||
ax.set_ylabel("Recovery error")
|
||||
ax.legend(fontsize=7, loc='upper left', framealpha=0.9)
|
||||
ax.grid(alpha=0.3)
|
||||
ax.set_aspect("equal")
|
||||
|
||||
fig.tight_layout()
|
||||
fig.savefig(os.path.join(args.out, "bound_verification.pdf"), bbox_inches="tight")
|
||||
print("Saved bound_verification.pdf")
|
||||
plt.close()
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════
|
||||
# Grid-specific plots (appendix)
|
||||
# ══════════════════════════════════════════════════════════════════════
|
||||
grid_dir = None
|
||||
for d in args.results_dirs:
|
||||
if "grid" in d:
|
||||
grid_dir = d
|
||||
break
|
||||
if grid_dir is None:
|
||||
print("No grid dir found, skipping decomposition and heatmaps.")
|
||||
return
|
||||
|
||||
grid_data = []
|
||||
for path in sorted(glob.glob(os.path.join(grid_dir, "*.json"))):
|
||||
with open(path) as f:
|
||||
grid_data.append(json.load(f))
|
||||
if not grid_data:
|
||||
print("No grid results."); return
|
||||
|
||||
errors_g = np.array([d["procrustes_mse"] for d in grid_data])
|
||||
epsilons = np.array([d["epsilon"] for d in grid_data])
|
||||
deltas = np.array([d["delta"] for d in grid_data])
|
||||
rhos = np.array([d["rho"] for d in grid_data])
|
||||
lambs = np.array([d["lamb"] for d in grid_data])
|
||||
|
||||
LAMB_MARKERS = {1e-6: 'h', 1e-5: 'H', 1e-4: 'p',
|
||||
1e-3: 'o', 5e-3: 's', 1e-2: 'D', 5e-2: '^', 1e-1: 'v', 5e-1: 'P'}
|
||||
RHO_MARKERS = {0.3: 'o', 0.5: 's', 0.7: 'D', 0.8: '^', 0.9: 'v', 0.95: 'P', 0.99: 'X'}
|
||||
|
||||
def scatter_by_lamb(ax, xvals, yvals):
|
||||
cmap = plt.cm.viridis
|
||||
norm = mpl.colors.Normalize(vmin=min(rhos), vmax=max(rhos))
|
||||
for lamb in sorted(set(lambs)):
|
||||
mask = lambs == lamb
|
||||
ax.scatter(xvals[mask], yvals[mask], c=rhos[mask], cmap=cmap, norm=norm,
|
||||
marker=LAMB_MARKERS.get(lamb, 'o'), s=30, alpha=0.8,
|
||||
edgecolors='k', linewidths=0.3, label=f"$\\lambda$={lamb:.0e}")
|
||||
return cmap, norm, r"Correlation [$\rho$]"
|
||||
|
||||
def scatter_by_rho(ax, xvals, yvals):
|
||||
log_lambs = np.log10(lambs)
|
||||
cmap = plt.cm.plasma
|
||||
norm = mpl.colors.Normalize(vmin=log_lambs.min(), vmax=log_lambs.max())
|
||||
for rho_val in sorted(set(rhos)):
|
||||
mask = rhos == rho_val
|
||||
ax.scatter(xvals[mask], yvals[mask], c=log_lambs[mask], cmap=cmap, norm=norm,
|
||||
marker=RHO_MARKERS.get(rho_val, 'o'), s=30, alpha=0.8,
|
||||
edgecolors='k', linewidths=0.3, label=f"$\\rho$={rho_val:.2f}")
|
||||
return cmap, norm, r"Regularization [$\log_{10}\lambda$]"
|
||||
|
||||
# ── Decomposition ──
|
||||
ylabel = r"$\min_{Q \in O(n)} \mathbb{E}[\|h(z) - Qz\|^2]$"
|
||||
fig, axes = plt.subplots(2, 2, figsize=(8, 6))
|
||||
x_configs = [
|
||||
(epsilons, r"$\varepsilon = \|\mathrm{Cov}(h(z)) - I\|_F$", r"Error vs $\varepsilon$"),
|
||||
(deltas, r"$\delta = \mathcal{L}(h) - 2(1{-}\rho)\,\mathrm{tr}(\Sigma)$", r"Error vs $\delta$"),
|
||||
]
|
||||
for row, scatter_fn in enumerate([scatter_by_lamb, scatter_by_rho]):
|
||||
for col, (xvals, xlabel, title) in enumerate(x_configs):
|
||||
ax = axes[row, col]
|
||||
cmap, norm, cbar_label = scatter_fn(ax, xvals, errors_g)
|
||||
ax.set_xlabel(xlabel); ax.set_ylabel(ylabel); ax.set_title(title)
|
||||
sm = plt.cm.ScalarMappable(cmap=cmap, norm=norm); sm.set_array([])
|
||||
plt.colorbar(sm, ax=ax, label=cbar_label, shrink=0.85)
|
||||
ax.legend(fontsize=5.5, loc='upper right', framealpha=0.9)
|
||||
ax.grid()
|
||||
fig.tight_layout()
|
||||
fig.savefig(os.path.join(args.out, "bound_decomposition.pdf"), bbox_inches="tight")
|
||||
print("Saved bound_decomposition.pdf")
|
||||
plt.close()
|
||||
|
||||
# ── Heatmaps ──
|
||||
unique_lambs = sorted(set(lambs))
|
||||
unique_rhos = sorted(set(rhos))
|
||||
for metric_key, title, cmap_name in [
|
||||
("r2_hz", "Linear $R^2$ (h -> z)", "viridis"),
|
||||
("orth_err_normalized", "Orth. error normalized", "viridis_r"),
|
||||
]:
|
||||
grid = np.full((len(unique_lambs), len(unique_rhos)), np.nan)
|
||||
counts = np.zeros_like(grid)
|
||||
for r in grid_data:
|
||||
li = unique_lambs.index(r["lamb"])
|
||||
ri = unique_rhos.index(r["rho"])
|
||||
val = r.get(metric_key)
|
||||
if val is not None:
|
||||
if np.isnan(grid[li, ri]):
|
||||
grid[li, ri] = 0
|
||||
grid[li, ri] += val
|
||||
counts[li, ri] += 1
|
||||
grid = np.where(counts > 0, grid / counts, np.nan)
|
||||
|
||||
fig, ax = plt.subplots(figsize=(7, 5))
|
||||
im = ax.imshow(grid, aspect="auto", origin="lower", cmap=cmap_name)
|
||||
plt.colorbar(im, ax=ax, label=title)
|
||||
ax.set_xticks(range(len(unique_rhos)))
|
||||
ax.set_xticklabels([f"{r:.2f}" for r in unique_rhos])
|
||||
ax.set_yticks(range(len(unique_lambs)))
|
||||
ax.set_yticklabels([f"{l:.0e}" for l in unique_lambs])
|
||||
ax.set_xlabel(r"$\rho$"); ax.set_ylabel(r"$\lambda$")
|
||||
ax.set_title(title)
|
||||
for i in range(len(unique_lambs)):
|
||||
for j in range(len(unique_rhos)):
|
||||
if not np.isnan(grid[i, j]):
|
||||
ax.text(j, i, f"{grid[i,j]:.3f}", ha="center", va="center", fontsize=6)
|
||||
fig.tight_layout()
|
||||
safe = metric_key.replace(".", "_")
|
||||
fig.savefig(os.path.join(args.out, f"heatmap_{safe}.pdf"), bbox_inches="tight")
|
||||
print(f"Saved heatmap_{safe}.pdf")
|
||||
plt.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,202 @@
|
||||
"""
|
||||
Generalized normal sweep across 4 mixings.
|
||||
|
||||
Emits three figures:
|
||||
1. fig_gennorm.pdf — 4-panel R^2(h -> z) vs alpha for SIGReg/VICReg/InfoNCE
|
||||
2. fig_gennorm_orth.pdf — 4-panel orthogonality error vs alpha (unconstrained ylim
|
||||
to show InfoNCE excursions off the chart)
|
||||
3. fig_gennorm_main.pdf — single-panel spiral-only headline figure for main text,
|
||||
matching the Fig.~4b style of the paper
|
||||
|
||||
Usage:
|
||||
python analysis/plot_gennorm.py --results_dir results/gennorm/ --out figures/
|
||||
"""
|
||||
import argparse, glob, json, os, re
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
from collections import defaultdict
|
||||
|
||||
MIXINGS = [("spiral", "Spiral"), ("banana", "Banana"),
|
||||
("sinusoid", "Sinusoid"), ("nvp", "NVP")]
|
||||
|
||||
MODES = ("lejepa", "whiten", "infonce")
|
||||
COLORS = {"lejepa": "#d62728", "whiten": "#1f77b4", "infonce": "#2ca02c"}
|
||||
LABELS = {"lejepa": "SIGReg", "whiten": "VICReg", "infonce": "InfoNCE"}
|
||||
|
||||
YLABELS = {
|
||||
"r2_hz_grid": r"Linear identifiability $R^2(h \to z)$",
|
||||
"r2_hz": r"Linear identifiability $R^2(h \to z)$",
|
||||
"orth_err_normalized_grid": r"Orthogonality error $\|\hat Q^\top \hat Q - I\|_F / \sqrt{n}$",
|
||||
"orth_err_normalized": r"Orthogonality error $\|\hat Q^\top \hat Q - I\|_F / \sqrt{n}$",
|
||||
}
|
||||
|
||||
# Sensible y-limits per metric. R^2 is bounded in [0,1] so we clip there.
|
||||
# Orthogonality error is unbounded above (Whitening/InfoNCE off-Gaussian can spike
|
||||
# into the tens), so we use log scale and let matplotlib autoscale.
|
||||
YLIMS = {
|
||||
"r2_hz_grid": (-0.05, 1.05),
|
||||
"r2_hz": (-0.05, 1.05),
|
||||
"orth_err_normalized_grid": None, # autoscale; log scale (see YSCALES) handles outliers
|
||||
"orth_err_normalized": None,
|
||||
}
|
||||
|
||||
# Y-axis scale per metric. Linear by default; log for orth error to compress
|
||||
# off-Gaussian excursions while still showing structure near zero.
|
||||
YSCALES = {
|
||||
"r2_hz_grid": "linear",
|
||||
"r2_hz": "linear",
|
||||
"orth_err_normalized_grid": "log",
|
||||
"orth_err_normalized": "log",
|
||||
}
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Data loading
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def load_groups(results_dir, metric):
|
||||
"""groups[(mixing, mode, alpha)] -> list of seed values for `metric`."""
|
||||
groups = defaultdict(list)
|
||||
for path in sorted(glob.glob(os.path.join(results_dir, "*.json"))):
|
||||
with open(path) as f:
|
||||
r = json.load(f)
|
||||
if r.get("experiment") != "gennorm":
|
||||
continue
|
||||
alpha = r.get("source_alpha")
|
||||
if alpha is None:
|
||||
m = re.search(r"alpha=([\d.]+)", r.get("run_name", ""))
|
||||
if m:
|
||||
alpha = float(m.group(1))
|
||||
if alpha is None or metric not in r:
|
||||
continue
|
||||
groups[(r["mixing"], r["mode"], alpha)].append(r[metric])
|
||||
return groups
|
||||
|
||||
|
||||
def curve(groups, mixing, mode):
|
||||
alphas = sorted({a for (mx, m, a) in groups if mx == mixing and m == mode})
|
||||
mu = np.array([np.mean(groups[(mixing, mode, a)]) for a in alphas])
|
||||
sd = np.array([np.std (groups[(mixing, mode, a)]) for a in alphas])
|
||||
return np.array(alphas), mu, sd
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Figure 1 & 2: 4-panel grids (one per metric)
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def plot_grid(groups, metric, out_path):
|
||||
fig, axes = plt.subplots(1, 4, figsize=(13, 3.0), sharey=True)
|
||||
use_log = YSCALES[metric] == "log"
|
||||
|
||||
for ax, (mix_key, mix_name) in zip(axes, MIXINGS):
|
||||
for mode in MODES:
|
||||
alphas, mu, sd = curve(groups, mix_key, mode)
|
||||
if len(alphas) == 0:
|
||||
continue
|
||||
ax.plot(alphas, mu, marker="o", ms=5, lw=1.8,
|
||||
color=COLORS[mode], label=LABELS[mode], zorder=3)
|
||||
# On log axes, clip the lower edge of the band away from zero
|
||||
# so fill_between doesn't disappear / warn.
|
||||
lower = mu - sd
|
||||
if use_log:
|
||||
lower = np.maximum(lower, 1e-3)
|
||||
ax.fill_between(alphas, lower, mu + sd,
|
||||
color=COLORS[mode], alpha=0.2, zorder=2)
|
||||
ax.set_xscale("log", base=2)
|
||||
if use_log:
|
||||
ax.set_yscale("log")
|
||||
ax.axvline(2.0, color="black", lw=0.7, ls="--", alpha=0.6, zorder=1)
|
||||
ax.set_xlabel(r"Source shape $\alpha$")
|
||||
ax.set_title(mix_name)
|
||||
if YLIMS[metric] is not None:
|
||||
ax.set_ylim(*YLIMS[metric])
|
||||
ax.grid(alpha=0.3, which="both" if use_log else "major")
|
||||
|
||||
axes[0].set_ylabel(YLABELS[metric])
|
||||
axes[-1].legend(frameon=False, loc="best")
|
||||
|
||||
fig.tight_layout()
|
||||
fig.savefig(out_path, bbox_inches="tight")
|
||||
print(f"Saved {out_path}")
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Figure 3: main-text single-panel headline (spiral mixing, all three methods)
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def plot_main_panel(groups, out_path):
|
||||
"""Single-panel spiral-only figure to sit next to the bound-verification panel
|
||||
in the main-text composite figure (Fig.~4b in the paper)."""
|
||||
FIGSIZE = 0.8 * np.array((4.0, 3.0))
|
||||
LW = 2.0
|
||||
MS = 6
|
||||
FONTSIZE = 11
|
||||
|
||||
rc_saved = plt.rcParams.copy()
|
||||
plt.rcParams.update({
|
||||
"font.size": FONTSIZE,
|
||||
"axes.labelsize": FONTSIZE,
|
||||
"xtick.labelsize": FONTSIZE - 1,
|
||||
"ytick.labelsize": FONTSIZE - 1,
|
||||
"legend.fontsize": FONTSIZE - 1,
|
||||
"axes.spines.top": False,
|
||||
"axes.spines.right": False,
|
||||
})
|
||||
|
||||
fig, ax = plt.subplots(figsize=FIGSIZE)
|
||||
|
||||
for mode in MODES:
|
||||
alphas, mu, sd = curve(groups, "spiral", mode)
|
||||
if len(alphas) == 0:
|
||||
continue
|
||||
ax.plot(alphas, mu, marker="o", ms=MS, lw=LW,
|
||||
color=COLORS[mode], label=LABELS[mode], zorder=3)
|
||||
ax.fill_between(alphas, mu - sd, mu + sd,
|
||||
color=COLORS[mode], alpha=0.2, zorder=2)
|
||||
|
||||
# Reference lines for canonical distributions
|
||||
ax.axvline(1.0, 0, 0.95, color="gray", lw=0.8, ls=":", alpha=0.7, zorder=1)
|
||||
ax.axvline(2.0, 0, 0.90, color="black", lw=0.8, ls="--", alpha=0.7, zorder=1)
|
||||
ax.text(1.0, 1.04, "Laplace", ha="center", va="bottom",
|
||||
fontsize=FONTSIZE - 1, color="gray")
|
||||
|
||||
ax.set_xscale("log", base=2)
|
||||
ax.set_ylabel(r"Linearity")
|
||||
ax.set_ylim(-0.05, 1.12)
|
||||
ax.grid(alpha=0.3, which="both")
|
||||
ax.legend(frameon=False, loc="lower right")
|
||||
ax.set_xticks([2**(-2), 2, 16])
|
||||
ax.set_xticklabels([r"$\leftarrow$ sparse", "Gaussian", r"uniform $\rightarrow$"])
|
||||
|
||||
fig.tight_layout()
|
||||
fig.savefig(out_path, bbox_inches="tight", dpi=500)
|
||||
print(f"Saved {out_path}")
|
||||
plt.close(fig)
|
||||
plt.rcParams.update(rc_saved)
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
# Driver
|
||||
# ──────────────────────────────────────────────────────────────────────────
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dir", default="results/gennorm/")
|
||||
p.add_argument("--out", default="figures/")
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
# 4-panel R^2 grid (appendix)
|
||||
groups_r2 = load_groups(args.results_dir, "r2_hz_grid")
|
||||
plot_grid(groups_r2, "r2_hz_grid",
|
||||
os.path.join(args.out, "fig_gennorm.pdf"))
|
||||
|
||||
# 4-panel orthogonality grid (appendix, autoscaled)
|
||||
groups_orth = load_groups(args.results_dir, "orth_err_normalized_grid")
|
||||
plot_grid(groups_orth, "orth_err_normalized_grid",
|
||||
os.path.join(args.out, "fig_gennorm_orth.pdf"))
|
||||
|
||||
# Single-panel main-text headline (spiral, all three methods)
|
||||
plot_main_panel(groups_r2,
|
||||
os.path.join(args.out, "fig_gennorm_main.pdf"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,73 @@
|
||||
"""
|
||||
Scaling plots: R² and orthogonality vs latent dimension N.
|
||||
|
||||
Usage:
|
||||
python analysis/plot_scaling.py --results_dir results/scaling/ --out figures/
|
||||
"""
|
||||
|
||||
import argparse, os, glob, json
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
|
||||
def load_results(results_dir):
|
||||
rows = []
|
||||
for path in sorted(glob.glob(os.path.join(results_dir, "*.json"))):
|
||||
with open(path) as f:
|
||||
r = json.load(f)
|
||||
rows.append({k: r.get(k) for k in [
|
||||
"N", "seed", "r2_zx", "r2_xz", "r2_zh", "r2_hz",
|
||||
"orth_err", "orth_err_normalized", "final_loss",
|
||||
"final_align", "final_sigreg", "final_whiten",
|
||||
"epsilon", "delta", "approx_bound", "procrustes_mse",
|
||||
]})
|
||||
return pd.DataFrame(rows)
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dir", default="results/scaling/")
|
||||
p.add_argument("--out", default="figures/")
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
df = load_results(args.results_dir)
|
||||
if len(df) == 0:
|
||||
print("No results."); return
|
||||
|
||||
summary = df.groupby("N").agg(
|
||||
r2_xz_mean=("r2_xz", "mean"), r2_xz_std=("r2_xz", "std"),
|
||||
r2_hz_mean=("r2_hz", "mean"), r2_hz_std=("r2_hz", "std"),
|
||||
orth_mean=("orth_err_normalized", "mean"), orth_std=("orth_err_normalized", "std"),
|
||||
).reset_index()
|
||||
dims = summary["N"].values
|
||||
|
||||
fig, axes = plt.subplots(1, 2, figsize=(8, 3.5))
|
||||
|
||||
ax = axes[0]
|
||||
ax.errorbar(dims, summary["r2_xz_mean"], yerr=summary["r2_xz_std"],
|
||||
fmt="o-", capsize=3, color="gray", label=r"Probe: $g(z) \to z$")
|
||||
ax.errorbar(dims, summary["r2_hz_mean"], yerr=summary["r2_hz_std"],
|
||||
fmt="s-", capsize=3, label=r"Probe: $f \circ g(z) \to z$")
|
||||
ax.set_xscale("log", base=2)
|
||||
ax.set_xlabel("Latent dimension $N$"); ax.set_ylabel(r"Linearity [$R^2$]")
|
||||
ax.set_title("Latent Recovery"); ax.set_ylim(-0.05, 1.05)
|
||||
ax.set_xticks(dims); ax.legend(); ax.grid(alpha=0.3)
|
||||
|
||||
ax = axes[1]
|
||||
ax.errorbar(dims, summary["orth_mean"], yerr=summary["orth_std"],
|
||||
fmt="D-", capsize=3, color="tab:green")
|
||||
ax.set_xscale("log", base=2)
|
||||
ax.set_xlabel("Latent dimension $N$")
|
||||
ax.set_ylabel(r"$\|A^\top A - I\|_F / \sqrt{N}$")
|
||||
ax.set_title("Orthogonality Error"); ax.set_xticks(dims); ax.grid(alpha=0.3)
|
||||
|
||||
fig.tight_layout()
|
||||
fig.savefig(os.path.join(args.out, "fig_scaling.pdf"), bbox_inches="tight")
|
||||
print("Saved fig_scaling.pdf")
|
||||
plt.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
Cross-experiment scatter plots (2x2 panel).
|
||||
|
||||
Usage:
|
||||
python analysis/plot_scatter.py --results_dirs results/2d results/scaling results/grid results/ablation --out figures/
|
||||
"""
|
||||
|
||||
import argparse, os, glob, json
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
EXPERIMENT_COLORS = {"2d": "tab:blue", "grid": "tab:red", "scaling": "tab:green", "ablation": "tab:orange"}
|
||||
EXPERIMENT_ORDER = ["grid", "scaling", "2d", "ablation"]
|
||||
|
||||
|
||||
def load_all(dirs):
|
||||
rows = []
|
||||
for d in dirs:
|
||||
for path in sorted(glob.glob(os.path.join(d, "*.json"))):
|
||||
try:
|
||||
with open(path) as f:
|
||||
rows.append(json.load(f))
|
||||
except Exception:
|
||||
pass
|
||||
return rows
|
||||
|
||||
|
||||
def scatter_by_experiment(ax, rows, x_key, y_key):
|
||||
for exp in EXPERIMENT_ORDER:
|
||||
pts = [r for r in rows if r.get("experiment") == exp
|
||||
and r.get(x_key) is not None and r.get(y_key) is not None]
|
||||
if not pts:
|
||||
continue
|
||||
ax.scatter([r[x_key] for r in pts],
|
||||
[r[y_key] for r in pts],
|
||||
c=EXPERIMENT_COLORS[exp],
|
||||
s=25, alpha=0.7, edgecolors='k', linewidths=0.3,
|
||||
label=exp, zorder=3)
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--results_dirs", nargs="+", required=True)
|
||||
p.add_argument("--out", default="figures/")
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
rows = load_all(args.results_dirs)
|
||||
if not rows:
|
||||
print("No results."); return
|
||||
|
||||
xlim = (5e-3, 1e0)
|
||||
ylim = (0.9, 1.01)
|
||||
|
||||
fig = plt.figure(figsize=0.65 * np.array((8, 7)))
|
||||
|
||||
# ── Total loss vs R² ──
|
||||
ax = plt.subplot(2, 2, 1)
|
||||
scatter_by_experiment(ax, rows, "final_loss", "r2_hz")
|
||||
ax.set_xlabel("Total loss")
|
||||
ax.set_ylabel("Linear Identifiability")
|
||||
ax.legend(fontsize=8)
|
||||
ax.grid(alpha=0.3)
|
||||
ax.set_xlim(*xlim)
|
||||
ax.set_ylim(*ylim)
|
||||
ax.set_xscale("log")
|
||||
|
||||
# ── Alignment vs R² ──
|
||||
ax = plt.subplot(2, 2, 2)
|
||||
scatter_by_experiment(ax, rows, "final_align", "r2_hz")
|
||||
ax.set_xlabel("Alignment loss")
|
||||
ax.set_ylabel("Linear Identifiability")
|
||||
ax.grid(alpha=0.3)
|
||||
ax.set_xlim(*xlim)
|
||||
ax.set_ylim(*ylim)
|
||||
ax.set_xscale("log")
|
||||
|
||||
# ── SIGReg vs R² ──
|
||||
ax = plt.subplot(2, 2, 3)
|
||||
scatter_by_experiment(ax, rows, "final_sigreg", "r2_hz")
|
||||
ax.set_xlabel("SIGReg loss")
|
||||
ax.set_ylabel("Linear Identifiability")
|
||||
ax.grid(alpha=0.3)
|
||||
ax.set_ylim(*ylim)
|
||||
ax.set_xscale("log")
|
||||
|
||||
# ── SIGReg vs whitening ──
|
||||
ax = plt.subplot(2, 2, 4)
|
||||
scatter_by_experiment(ax, rows, "final_sigreg", "final_whiten")
|
||||
ax.set_xlabel("SIGReg loss")
|
||||
ax.set_ylabel("Whitening loss")
|
||||
ax.grid(alpha=0.3)
|
||||
ax.set_xscale("log")
|
||||
ax.set_yscale("log")
|
||||
|
||||
fig.tight_layout()
|
||||
fig.savefig(os.path.join(args.out, "scatter_plots.pdf"), bbox_inches="tight")
|
||||
print("Saved scatter_plots.pdf")
|
||||
plt.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,10 @@
|
||||
#!/bin/bash
|
||||
|
||||
python analysis/aggregate.py --results_dir results/2d/ --out results/2d/summary.csv
|
||||
python analysis/aggregate.py --results_dir results/ --recursive --out results/all.csv
|
||||
python analysis/plot_2d.py --results_dir results/2d/ --out figures/
|
||||
python analysis/plot_bound.py --results_dirs results/grid results/2d results/scaling results/ablation --out figures/
|
||||
python analysis/plot_ablation.py --results_dir results/ablation/ --out figures/
|
||||
python analysis/plot_scaling.py --results_dir results/scaling/ --out figures/
|
||||
python analysis/plot_scatter.py --results_dirs results/2d results/scaling results/grid results/ablation --out figures/
|
||||
python analysis/make_table_scaling.py --results_dir results/scaling/
|
||||
@@ -0,0 +1,30 @@
|
||||
# 2D illustration + Gaussian half of the regularizer ablation
|
||||
# python run.py --config configs/2d.yaml --run spiral_lejepa --seed 1337
|
||||
|
||||
experiment: 2d
|
||||
out: results/2d
|
||||
|
||||
# Data
|
||||
N: 2
|
||||
source_dist: gaussian
|
||||
num_eval: 10000
|
||||
|
||||
# Training (shared)
|
||||
steps: 20000
|
||||
lr: 3.0e-3
|
||||
batch_size: 256
|
||||
rho: 0.95
|
||||
log_every: 500
|
||||
|
||||
# Per-run specs: 4 mixings x 2 objectives
|
||||
runs:
|
||||
spiral_lejepa: {mixing: spiral, encoder: mlp, hidden: 256, mode: lejepa, lamb: 1.0e-3}
|
||||
spiral_whiten: {mixing: spiral, encoder: mlp, hidden: 256, mode: whiten, lamb: 0.5}
|
||||
banana_lejepa: {mixing: banana, encoder: mlp, hidden: 256, mode: lejepa, lamb: 1.0e-3}
|
||||
banana_whiten: {mixing: banana, encoder: mlp, hidden: 256, mode: whiten, lamb: 0.5}
|
||||
sinusoid_lejepa: {mixing: sinusoid, encoder: mlp, hidden: 256, mode: lejepa, lamb: 1.0e-3}
|
||||
sinusoid_whiten: {mixing: sinusoid, encoder: mlp, hidden: 256, mode: whiten, lamb: 0.5}
|
||||
nvp_lejepa: {mixing: nvp, encoder: matched, n_layers: 8, mode: lejepa, lamb: 1.0e-3}
|
||||
nvp_whiten: {mixing: nvp, encoder: matched, n_layers: 8, mode: whiten, lamb: 0.5}
|
||||
|
||||
seeds: [1337, 1338, 1339]
|
||||
@@ -0,0 +1,32 @@
|
||||
# Generalized normal sweep across mixings (main-text figure)
|
||||
# python run.py --config configs/gennorm.yaml --run spiral_lejepa --alpha 2.0 --seed 1337
|
||||
|
||||
experiment: gennorm
|
||||
out: results/gennorm
|
||||
|
||||
N: 2
|
||||
source_dist: gennorm # alpha provided per-run via CLI
|
||||
num_eval: 10000
|
||||
|
||||
steps: 20000
|
||||
lr: 3.0e-3
|
||||
batch_size: 256
|
||||
rho: 0.95
|
||||
log_every: 500
|
||||
|
||||
runs:
|
||||
spiral_lejepa: {mixing: spiral, encoder: mlp, hidden: 256, mode: lejepa, lamb: 1.0e-3}
|
||||
spiral_whiten: {mixing: spiral, encoder: mlp, hidden: 256, mode: whiten, lamb: 0.5}
|
||||
spiral_infonce: {mixing: spiral, encoder: mlp, hidden: 256, mode: infonce, sigma: 1.0}
|
||||
banana_lejepa: {mixing: banana, encoder: mlp, hidden: 256, mode: lejepa, lamb: 1.0e-3}
|
||||
banana_whiten: {mixing: banana, encoder: mlp, hidden: 256, mode: whiten, lamb: 0.5}
|
||||
banana_infonce: {mixing: banana, encoder: mlp, hidden: 256, mode: infonce, sigma: 1.0}
|
||||
sinusoid_lejepa: {mixing: sinusoid, encoder: mlp, hidden: 256, mode: lejepa, lamb: 1.0e-3}
|
||||
sinusoid_whiten: {mixing: sinusoid, encoder: mlp, hidden: 256, mode: whiten, lamb: 0.5}
|
||||
sinusoid_infonce: {mixing: sinusoid, encoder: mlp, hidden: 256, mode: infonce, sigma: 1.0}
|
||||
nvp_lejepa: {mixing: nvp, encoder: matched, n_layers: 8, mode: lejepa, lamb: 1.0e-3}
|
||||
nvp_whiten: {mixing: nvp, encoder: matched, n_layers: 8, mode: whiten, lamb: 0.5}
|
||||
nvp_infonce: {mixing: nvp, encoder: matched, n_layers: 8, mode: infonce, sigma: 1.0}
|
||||
|
||||
alphas: [0.125, 0.25, 0.5, 1.0, 2.0, 4.0, 8.0, 16.0, 32.0]
|
||||
seeds: [1337, 1338, 1339]
|
||||
@@ -0,0 +1,27 @@
|
||||
# Grid search over lambda and rho (bound verification figure)
|
||||
# python run.py --config configs/grid.yaml --lamb 0.01 --rho 0.9 --seed 0
|
||||
|
||||
experiment: grid
|
||||
out: results/grid
|
||||
|
||||
# Data
|
||||
N: 2
|
||||
source_dist: gaussian
|
||||
num_eval: 10000
|
||||
|
||||
# Training (shared)
|
||||
steps: 20000
|
||||
lr: 3.0e-3
|
||||
batch_size: 256
|
||||
log_every: 500
|
||||
|
||||
# Encoder
|
||||
encoder: mlp
|
||||
hidden: 256
|
||||
mixing: spiral
|
||||
mode: lejepa
|
||||
|
||||
# Sweep dimensions
|
||||
lambs: [1.0e-6, 1.0e-5, 1.0e-4, 1.0e-3, 5.0e-3, 1.0e-2, 5.0e-2, 1.0e-1, 5.0e-1]
|
||||
rhos: [0.3, 0.5, 0.7, 0.8, 0.9, 0.95, 0.99]
|
||||
seeds: [0, 1, 2]
|
||||
@@ -0,0 +1,29 @@
|
||||
# Reacher pixel-observation identifiability experiment
|
||||
# Works for both OU and trajectory data — just point --data_dir at the right place.
|
||||
#
|
||||
# Prerender:
|
||||
# python prerender.py eval
|
||||
# python prerender.py ou --rho 0.95
|
||||
# python prerender.py traj --delta 16 --h5_path data/reacher.h5
|
||||
#
|
||||
# Train:
|
||||
# python run_reacher.py --config configs/reacher.yaml \
|
||||
# --data_dir data/reacher/ou/rho=0.95
|
||||
|
||||
experiment: reacher
|
||||
out: results/reacher
|
||||
data_root: data/reacher
|
||||
|
||||
# Model
|
||||
d_latent: 2
|
||||
|
||||
# Training
|
||||
epochs: 100
|
||||
batch_size: 256
|
||||
lr: 3.0e-3
|
||||
n_slices: 256
|
||||
n_eval_fast: 2000
|
||||
|
||||
# Sweep dimensions
|
||||
lambs: [1.0e-3, 5.0e-3, 1.0e-2, 5.0e-2]
|
||||
seeds: [0, 1, 2]
|
||||
@@ -0,0 +1,36 @@
|
||||
# Scaling experiment (paper figure)
|
||||
# python run.py --config configs/scaling.yaml --N 16 --seed 0
|
||||
# python run.py --config configs/scaling.yaml --N 16 --seed 0 --mode infonce
|
||||
# python run.py --config configs/scaling.yaml --N 16 --seed 0 --mode whiten
|
||||
|
||||
experiment: scaling
|
||||
out: results/scaling
|
||||
|
||||
# Data
|
||||
source_dist: gaussian
|
||||
num_eval: 10000
|
||||
|
||||
# Training (shared)
|
||||
steps: 20000
|
||||
lr: 3.0e-3
|
||||
batch_size: 256
|
||||
rho: 0.95
|
||||
log_every: 500
|
||||
|
||||
# Encoder
|
||||
encoder: matched
|
||||
n_layers: 4
|
||||
mode: lejepa # default; override with --mode
|
||||
|
||||
# Mode-specific defaults (used based on --mode)
|
||||
lamb: 1.0e-6 # for lejepa
|
||||
lamb_whiten: 0.5 # used when mode=whiten
|
||||
sigma: 1.0 # for infonce
|
||||
|
||||
# Mixing
|
||||
mixing: coupling
|
||||
|
||||
# Sweep dimensions
|
||||
dims: [2, 4, 8, 16, 32, 64, 128, 256, 512, 1024]
|
||||
seeds: [0, 1, 2, 3, 4]
|
||||
K: 3 # parallel encoder runs per (N, seed); pick lowest loss
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Data generation: latent sources and OU augmentation."""
|
||||
|
||||
import math
|
||||
import torch
|
||||
|
||||
|
||||
def _gennorm_unit_var_scale(alpha):
|
||||
"""Scale β so gennorm(α, β) has unit variance: β = sqrt(Γ(1/α) / Γ(3/α))."""
|
||||
return math.exp(0.5 * (math.lgamma(1.0 / alpha) - math.lgamma(3.0 / alpha)))
|
||||
|
||||
|
||||
def sample_latents(D, N, dist="gaussian", device="cuda", alpha=None):
|
||||
"""Sample D points in R^N (unit variance)."""
|
||||
if dist == "gaussian":
|
||||
return torch.randn(D, N, device=device)
|
||||
elif dist == "laplace":
|
||||
return torch.distributions.Laplace(0, 1 / (2 ** 0.5)).sample((D, N)).to(device)
|
||||
elif dist == "gennorm":
|
||||
if alpha is None:
|
||||
raise ValueError("gennorm requires alpha")
|
||||
scale = _gennorm_unit_var_scale(alpha)
|
||||
u = torch.distributions.Gamma(1.0 / alpha, 1.0).sample((D, N)).to(device)
|
||||
sign = torch.randint(0, 2, (D, N), device=device).float() * 2 - 1
|
||||
return scale * sign * u.pow(1.0 / alpha)
|
||||
else:
|
||||
raise ValueError(f"Unknown distribution: {dist}")
|
||||
|
||||
|
||||
def ou_augment(z, rho, n_views=2, dist="gaussian", alpha=None):
|
||||
"""OU channel: z' = ρz + √(1-ρ²)η, η drawn from same dist as source.
|
||||
Returns (V, B, N)."""
|
||||
fac = (1 - rho ** 2) ** 0.5
|
||||
D, N = z.shape
|
||||
eta = sample_latents(n_views * D, N, dist=dist, device=z.device, alpha=alpha)
|
||||
eta = eta.reshape(n_views, D, N)
|
||||
return rho * z.unsqueeze(0) + fac * eta
|
||||
@@ -0,0 +1,143 @@
|
||||
"""
|
||||
Core training engine — single function used by all experiments.
|
||||
|
||||
Handles: LR schedule (warmup + cosine), online data generation,
|
||||
loss computation (lejepa or whiten), periodic evaluation of ALL metrics
|
||||
on a fixed eval set, standardized output schema.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from .losses import SIGReg, whitening_loss, alignment_loss, infonce_loss
|
||||
from .data import sample_latents, ou_augment
|
||||
from .metrics import compute_all_metrics
|
||||
|
||||
|
||||
def warmup_cosine_lr(step, total_steps, base_lr):
|
||||
"""Constant for first half, cosine decay for second half."""
|
||||
warmup = total_steps // 2
|
||||
if step < warmup:
|
||||
return base_lr
|
||||
t = (step - warmup) / (total_steps - warmup)
|
||||
return base_lr * 0.5 * (1 + np.cos(np.pi * t))
|
||||
|
||||
|
||||
def train_and_evaluate(
|
||||
encoder,
|
||||
mix_fn,
|
||||
*,
|
||||
N,
|
||||
rho,
|
||||
lamb,
|
||||
sigma=1.0,
|
||||
mode="lejepa",
|
||||
source_dist="gaussian",
|
||||
source_alpha=None,
|
||||
steps=20000,
|
||||
batch_size=256,
|
||||
lr=3e-3,
|
||||
z_eval,
|
||||
log_every=100,
|
||||
device="cuda",
|
||||
):
|
||||
"""Train encoder and evaluate periodically.
|
||||
|
||||
Args:
|
||||
encoder: nn.Module, x -> h
|
||||
mix_fn: callable, z -> x
|
||||
N: latent dimension
|
||||
rho: OU correlation
|
||||
lamb: regularization weight
|
||||
mode: "lejepa" or "whiten"
|
||||
source_dist: "gaussian", "laplace", or "gennorm"
|
||||
steps: total training steps
|
||||
batch_size: batch size (online data)
|
||||
lr: peak learning rate
|
||||
z_eval: (num_eval, N) fixed eval tensor
|
||||
log_every: eval frequency
|
||||
device: torch device string
|
||||
|
||||
Returns:
|
||||
encoder: trained encoder
|
||||
log: dict of lists — training curves and periodic eval metrics
|
||||
"""
|
||||
sigreg = SIGReg().to(device)
|
||||
opt = torch.optim.AdamW(encoder.parameters(), lr=lr)
|
||||
|
||||
# Precompute eval mixing (constant across training)
|
||||
x_eval = mix_fn(z_eval)
|
||||
|
||||
log_keys = [
|
||||
"step", "lr",
|
||||
# training losses
|
||||
"align", "sigreg", "whiten", "total",
|
||||
# eval metrics
|
||||
"r2_zx", "r2_xz", "r2_zh", "r2_hz",
|
||||
"orth_err", "orth_err_normalized",
|
||||
"epsilon", "delta", "D_bound", "approx_bound",
|
||||
"procrustes_mse", "L_h", "trace_cov",
|
||||
]
|
||||
log = {k: [] for k in log_keys}
|
||||
|
||||
for step in range(steps + 1):
|
||||
# LR schedule
|
||||
current_lr = warmup_cosine_lr(step, steps, lr)
|
||||
for pg in opt.param_groups:
|
||||
pg["lr"] = current_lr
|
||||
|
||||
# Online data
|
||||
z_batch = sample_latents(batch_size, N, dist=source_dist,
|
||||
device=device, alpha=source_alpha)
|
||||
z_aug = ou_augment(z_batch, rho, dist=source_dist, alpha=source_alpha) # (2, B, N)
|
||||
h = encoder(mix_fn(z_aug).flatten(0, 1)).reshape(2, batch_size, N)
|
||||
|
||||
align = alignment_loss(h)
|
||||
sig = sigreg(h)
|
||||
wht = whitening_loss(h)
|
||||
|
||||
if mode == "lejepa":
|
||||
loss = lamb * sig + (1 - lamb) * align
|
||||
elif mode == "whiten":
|
||||
loss = lamb * wht + (1 - lamb) * align
|
||||
elif mode == "infonce":
|
||||
loss = infonce_loss(h, sigma)
|
||||
else:
|
||||
raise ValueError(f"Unknown mode: {mode}")
|
||||
|
||||
opt.zero_grad()
|
||||
loss.backward()
|
||||
opt.step()
|
||||
|
||||
if step % log_every == 0 or (step < 1000 and step % 100 == 0):
|
||||
log["step"].append(step)
|
||||
log["lr"].append(current_lr)
|
||||
log["align"].append(align.item())
|
||||
log["sigreg"].append(sig.item())
|
||||
log["whiten"].append(wht.item())
|
||||
log["total"].append(loss.item())
|
||||
|
||||
# Full eval on fixed set
|
||||
encoder.eval()
|
||||
with torch.no_grad():
|
||||
h_eval = encoder(x_eval)
|
||||
z_prime = ou_augment(
|
||||
z_eval, rho, n_views=1,
|
||||
dist=source_dist, alpha=source_alpha
|
||||
).squeeze(0)
|
||||
h_prime = encoder(mix_fn(z_prime))
|
||||
|
||||
metrics = compute_all_metrics(z_eval, x_eval, h_eval, h_prime, rho, N)
|
||||
|
||||
for k, v in metrics.items():
|
||||
log[k].append(v)
|
||||
|
||||
encoder.train()
|
||||
|
||||
if step % (log_every * 10) == 0:
|
||||
print(f" step {step:5d} | lr={current_lr:.1e} "
|
||||
f"align={align.item():.2e} sig={sig.item():.1f} "
|
||||
f"R²(h->z)={metrics['r2_hz']:.4f} "
|
||||
f"orth={metrics['orth_err']:.4f}")
|
||||
|
||||
return encoder, log
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Loss functions."""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class SIGReg(nn.Module):
|
||||
"""Sliced characteristic function regularizer (Balestriero & LeCun 2025)."""
|
||||
|
||||
def __init__(self, knots=17, n_slices=256, t_max=3.0):
|
||||
super().__init__()
|
||||
self.n_slices = n_slices
|
||||
t = torch.linspace(0, t_max, knots)
|
||||
dt = t_max / (knots - 1)
|
||||
w = torch.full((knots,), 2 * dt)
|
||||
w[[0, -1]] = dt
|
||||
self.register_buffer("t", t)
|
||||
self.register_buffer("phi", torch.exp(-t**2 / 2))
|
||||
self.register_buffer("weights", w * torch.exp(-t**2 / 2))
|
||||
|
||||
def forward(self, h):
|
||||
"""h: (V, B, N) -> scalar."""
|
||||
flat = h.flatten(0, 1)
|
||||
A = F.normalize(torch.randn(flat.size(-1), self.n_slices, device=flat.device), 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)
|
||||
|
||||
|
||||
def whitening_loss(h):
|
||||
"""||Cov(h) - I||²_F. h: (V, B, N) -> scalar."""
|
||||
flat = h.flatten(0, 1)
|
||||
flat = flat - flat.mean(dim=0)
|
||||
cov = (flat.T @ flat) / (flat.shape[0] - 1)
|
||||
return (cov - torch.eye(flat.shape[1], device=h.device)).square().mean()
|
||||
|
||||
|
||||
def alignment_loss(h):
|
||||
"""Pull positive-pair views together. h: (V, B, N) -> scalar."""
|
||||
return (h.mean(0) - h).square().mean()
|
||||
|
||||
|
||||
def infonce_loss(h, sigma):
|
||||
"""Symmetric Gaussian-kernel InfoNCE: sim(u, v) = -||u - v||² / (2σ²).
|
||||
h: (V, B, N) with V=2 views. Negatives are other batch elements.
|
||||
"""
|
||||
h1, h2 = h[0], h[1] # (B, N) each
|
||||
d12 = ((h1.unsqueeze(1) - h2.unsqueeze(0)) ** 2).sum(-1) # (B, B)
|
||||
sim = -d12 / (2 * sigma ** 2)
|
||||
loss_a = -(sim.diag() - torch.logsumexp(sim, dim=1)).mean()
|
||||
loss_b = -(sim.diag() - torch.logsumexp(sim, dim=0)).mean()
|
||||
return 0.5 * (loss_a + loss_b)
|
||||
@@ -0,0 +1,65 @@
|
||||
"""Evaluation metrics — standardized across all experiments."""
|
||||
|
||||
import torch
|
||||
|
||||
def bidirectional_r2(a, b):
|
||||
"""R²(a->b) and R²(b->a) via torch lstsq on GPU. a, b are tensors."""
|
||||
def _r2(x, y):
|
||||
x1 = torch.cat([x, torch.ones(len(x), 1, device=x.device)], dim=1)
|
||||
W = torch.linalg.lstsq(x1, y).solution
|
||||
ss_res = ((y - x1 @ W) ** 2).sum()
|
||||
ss_tot = ((y - y.mean(0)) ** 2).sum()
|
||||
return (1 - ss_res / ss_tot).item()
|
||||
return _r2(a, b), _r2(b, a)
|
||||
|
||||
|
||||
def compute_all_metrics(z, x, h, h_prime, rho, N):
|
||||
"""All metrics on GPU. z, x, h, h_prime are torch tensors."""
|
||||
r2_zx, r2_xz = bidirectional_r2(z, x)
|
||||
r2_zh, r2_hz = bidirectional_r2(z, h)
|
||||
|
||||
# Orthogonality
|
||||
z1 = torch.cat([z, torch.ones(len(z), 1, device=z.device)], dim=1)
|
||||
W = torch.linalg.lstsq(z1, h).solution
|
||||
A = W[:N].T
|
||||
orth_err = torch.linalg.norm(A.T @ A - torch.eye(N, device=A.device), 'fro').item()
|
||||
orth_err_normalized = orth_err / (N ** 0.5)
|
||||
|
||||
# Bound quantities
|
||||
cov_h = torch.cov(h.T)
|
||||
epsilon = torch.linalg.norm(cov_h - torch.eye(N, device=h.device), 'fro').item()
|
||||
trace_cov = torch.trace(cov_h).item()
|
||||
L_h = ((h_prime - h) ** 2).sum(dim=1).mean().item()
|
||||
delta = max(L_h - 2 * (1 - rho) * trace_cov, 0.0)
|
||||
spectral_gap = 2 * rho * (1 - rho)
|
||||
D_bound = delta / spectral_gap if spectral_gap > 0 else float("inf")
|
||||
approx_bound = D_bound + (epsilon + D_bound) ** 2
|
||||
|
||||
# Procrustes
|
||||
M = (h.T @ z) / len(z)
|
||||
U, S, Vt = torch.linalg.svd(M)
|
||||
Q = U @ Vt
|
||||
procrustes_mse = ((h - z @ Q.T) ** 2).sum(dim=1).mean().item()
|
||||
|
||||
return {
|
||||
"r2_zx": r2_zx, "r2_xz": r2_xz,
|
||||
"r2_zh": r2_zh, "r2_hz": r2_hz,
|
||||
"orth_err": orth_err, "orth_err_normalized": orth_err_normalized,
|
||||
"epsilon": epsilon, "delta": delta, "D_bound": D_bound,
|
||||
"approx_bound": approx_bound, "procrustes_mse": procrustes_mse,
|
||||
"L_h": L_h, "trace_cov": trace_cov,
|
||||
}
|
||||
|
||||
|
||||
def compute_recovery_metrics(z, h, N, suffix=""):
|
||||
"""R² in both directions + orthogonality. Suffix appended to keys."""
|
||||
r2_zh, r2_hz = bidirectional_r2(z, h)
|
||||
z1 = torch.cat([z, torch.ones(len(z), 1, device=z.device)], dim=1)
|
||||
W = torch.linalg.lstsq(z1, h).solution
|
||||
A = W[:N].T
|
||||
orth_err = torch.linalg.norm(A.T @ A - torch.eye(N, device=A.device), 'fro').item()
|
||||
return {
|
||||
f"r2_zh{suffix}": r2_zh, f"r2_hz{suffix}": r2_hz,
|
||||
f"orth_err{suffix}": orth_err,
|
||||
f"orth_err_normalized{suffix}": orth_err / (N ** 0.5),
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
"""Nonlinear mixing functions."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
# ── 2D mixing functions ──────────────────────────────────────────────────────
|
||||
|
||||
def mix_spiral(z):
|
||||
"""g(z) = R(π‖z‖) z — measure-preserving spiral diffeomorphism."""
|
||||
norms = z.norm(dim=-1) * torch.pi
|
||||
c, s = norms.cos(), norms.sin()
|
||||
R = torch.stack([torch.stack([c, -s], dim=-1),
|
||||
torch.stack([s, c], dim=-1)], dim=-2)
|
||||
return (R @ z.unsqueeze(-1)).squeeze(-1)
|
||||
|
||||
|
||||
def mix_banana(z):
|
||||
"""Banana: x0 = z0, x1 = z1 + z0²."""
|
||||
return torch.stack([z[..., 0], z[..., 1] + z[..., 0] ** 2], dim=-1)
|
||||
|
||||
|
||||
def mix_sinusoid(z):
|
||||
"""Sinusoidal shear: x0 = z0 + sin(1.5 z1), x1 = z1."""
|
||||
return torch.stack([z[..., 0] + torch.sin(1.5 * z[..., 1]), z[..., 1]], dim=-1)
|
||||
|
||||
|
||||
MIXINGS_2D = {
|
||||
"spiral": mix_spiral,
|
||||
"banana": mix_banana,
|
||||
"sinusoid": mix_sinusoid,
|
||||
# "nvp" handled via make_coupling_mixing(N=2, n_layers=...)
|
||||
}
|
||||
|
||||
|
||||
# ── Coupling-layer mixing (any dimension) ────────────────────────────────────
|
||||
|
||||
def make_coupling_mixing(N, n_layers=4, seed=1337, device="cuda"):
|
||||
"""RealNVP-style coupling layers. Works for any even N (including N=2)."""
|
||||
half = N // 2
|
||||
torch.manual_seed(seed)
|
||||
Ws = []
|
||||
for _ in range(n_layers):
|
||||
W, _ = torch.linalg.qr(torch.randn(half, half, device=device))
|
||||
Ws.append(W * 2.0)
|
||||
|
||||
def mix(z):
|
||||
for i, W in enumerate(Ws):
|
||||
z1, z2 = z[..., :half], z[..., half:]
|
||||
if i % 2 == 0:
|
||||
z2 = z2 + torch.tanh(z1 @ W)
|
||||
else:
|
||||
z1 = z1 + torch.tanh(z2 @ W)
|
||||
z = torch.cat([z1, z2], dim=-1)
|
||||
return z
|
||||
|
||||
return mix
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Encoder architectures."""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
|
||||
|
||||
def make_mlp_encoder(N, hidden=256, n_layers=4, device="cuda"):
|
||||
"""MLP encoder."""
|
||||
layers = [nn.Linear(N, hidden), nn.GELU()]
|
||||
for _ in range(n_layers - 1):
|
||||
layers += [nn.Linear(hidden, hidden), nn.GELU()]
|
||||
layers.append(nn.Linear(hidden, N))
|
||||
return nn.Sequential(*layers).to(device)
|
||||
|
||||
|
||||
class MatchedEncoder(nn.Module):
|
||||
"""Inverse coupling-layer encoder matched to NVP mixing architecture."""
|
||||
|
||||
def __init__(self, N, n_layers=4, device="cuda"):
|
||||
super().__init__()
|
||||
half = N // 2
|
||||
self.half = half
|
||||
self.n_layers = n_layers
|
||||
self.Ws = nn.ParameterList([
|
||||
nn.Parameter(torch.randn(half, half, device=device) / np.sqrt(half))
|
||||
for _ in range(n_layers)
|
||||
])
|
||||
|
||||
def forward(self, x):
|
||||
for i, W in reversed(list(enumerate(self.Ws))):
|
||||
z1, z2 = x[..., :self.half], x[..., self.half:]
|
||||
if i % 2 == 0:
|
||||
z2 = z2 - torch.tanh(z1 @ W)
|
||||
else:
|
||||
z1 = z1 - torch.tanh(z2 @ W)
|
||||
x = torch.cat([z1, z2], dim=-1)
|
||||
return x
|
||||
|
||||
|
||||
def make_matched_encoder(N, n_layers=4, seed=42, device="cuda"):
|
||||
torch.manual_seed(seed)
|
||||
return MatchedEncoder(N, n_layers=n_layers, device=device).to(device)
|
||||
|
||||
|
||||
def make_cnn_encoder(d_latent=2, device="cuda"):
|
||||
return nn.Sequential(
|
||||
nn.Conv2d(3, 32, 4, 2, 1), nn.BatchNorm2d(32), nn.GELU(),
|
||||
nn.Conv2d(32, 64, 4, 2, 1), nn.BatchNorm2d(64), nn.GELU(),
|
||||
nn.Conv2d(64, 128, 4, 2, 1), nn.BatchNorm2d(128), nn.GELU(),
|
||||
nn.Conv2d(128, 256, 4, 2, 1), nn.BatchNorm2d(256), nn.GELU(),
|
||||
torch.nn.AvgPool2d(4), nn.Flatten(),
|
||||
nn.Linear(256, 256), nn.BatchNorm1d(256), nn.GELU(),
|
||||
nn.Linear(256, d_latent),
|
||||
).to(device)
|
||||
@@ -0,0 +1,95 @@
|
||||
"""DMC Reacher rendering and dataset utilities."""
|
||||
|
||||
import os
|
||||
os.environ.setdefault("MUJOCO_GL", "egl")
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from dm_control import suite
|
||||
from tqdm import tqdm
|
||||
|
||||
from .data import ou_augment
|
||||
|
||||
|
||||
def make_env():
|
||||
return suite.load(domain_name="reacher", task_name="hard")
|
||||
|
||||
|
||||
def render_at(env, qpos, target, height=64, width=64):
|
||||
"""Set joint angles and render → (3, H, W) float32 in [0, 1]."""
|
||||
env.physics.data.qpos[:2] = qpos
|
||||
env.physics.data.qvel[:] = 0
|
||||
env.physics.named.model.geom_pos['target', :2] = target
|
||||
env.physics.forward()
|
||||
rgb = env.physics.render(height=height, width=width, camera_id=0)
|
||||
return rgb.transpose(2, 0, 1).astype(np.float32) / 255.0
|
||||
|
||||
|
||||
def render_batch(env, qpos_batch, target, height=64, width=64):
|
||||
"""Render batch → (N, 3, H, W)."""
|
||||
N = len(qpos_batch)
|
||||
imgs = np.empty((N, 3, height, width), dtype=np.float32)
|
||||
for i in tqdm(range(N), desc="Rendering"):
|
||||
imgs[i] = render_at(env, qpos_batch[i], target, height, width)
|
||||
return imgs
|
||||
|
||||
|
||||
def generate_ou_image_pairs(env, N, rho, target, seed=9999):
|
||||
"""
|
||||
Sample OU latent pairs, render both → (img_t, img_tp1, z_t, z_tp1).
|
||||
|
||||
Uses the same OU process as the rest of the repo but renders through MuJoCo.
|
||||
"""
|
||||
rng = np.random.default_rng(seed)
|
||||
z_t = rng.standard_normal((N, 2)).astype(np.float32)
|
||||
eps = rng.standard_normal((N, 2)).astype(np.float32)
|
||||
z_tp1 = rho * z_t + np.sqrt(1 - rho**2) * eps
|
||||
|
||||
print(f"Rendering {2 * N} images (rho={rho})...")
|
||||
img_t = render_batch(env, z_t, target)
|
||||
img_tp1 = render_batch(env, z_tp1, target)
|
||||
return img_t, img_tp1, z_t, z_tp1
|
||||
|
||||
|
||||
def normalize_images(img_t, img_tp1, img_eval=None):
|
||||
"""Per-channel mean/std normalization. Returns normalized arrays + stats."""
|
||||
mean = img_t.mean(axis=(0, 2, 3), keepdims=True)
|
||||
std = img_t.std(axis=(0, 2, 3), keepdims=True) + 1e-6
|
||||
img_t = (img_t - mean) / std
|
||||
img_tp1 = (img_tp1 - mean) / std
|
||||
if img_eval is not None:
|
||||
img_eval = (img_eval - mean) / std
|
||||
return img_t, img_tp1, img_eval, mean, std
|
||||
return img_t, img_tp1, mean, std
|
||||
|
||||
|
||||
def solve_ik_grid(env, target, n_grid=200):
|
||||
"""Find joint angles that place fingertip at target via grid search."""
|
||||
best_dist, best_qpos = np.inf, None
|
||||
for q0 in np.linspace(-np.pi, np.pi, n_grid):
|
||||
for q1 in np.linspace(-np.pi, np.pi, n_grid):
|
||||
env.physics.data.qpos[:2] = [q0, q1]
|
||||
env.physics.named.model.geom_pos['target', :2] = target
|
||||
env.physics.forward()
|
||||
tip = env.physics.named.data.geom_xpos['finger'][:2]
|
||||
d = np.linalg.norm(tip - target)
|
||||
if d < best_dist:
|
||||
best_dist = d
|
||||
best_qpos = np.array([q0, q1])
|
||||
return best_qpos, best_dist
|
||||
|
||||
|
||||
class ReacherOUDataset(torch.utils.data.Dataset):
|
||||
"""Prerendered OU image pairs with ground-truth latents."""
|
||||
|
||||
def __init__(self, img_t, img_tp1, z_t, z_tp1):
|
||||
self.img_t = torch.from_numpy(img_t)
|
||||
self.img_tp1 = torch.from_numpy(img_tp1)
|
||||
self.z_t = torch.from_numpy(z_t)
|
||||
self.z_tp1 = torch.from_numpy(z_tp1)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.img_t)
|
||||
|
||||
def __getitem__(self, i):
|
||||
return self.img_t[i], self.img_tp1[i], self.z_t[i], self.z_tp1[i]
|
||||
@@ -0,0 +1,279 @@
|
||||
"""
|
||||
Pre-render all Reacher datasets to disk.
|
||||
|
||||
Three modes:
|
||||
eval — 10k Gaussian samples, rendered once, shared by all runs
|
||||
ou — 100k OU pairs for a given rho
|
||||
traj — 100k pairs subsampled from LeWM trajectories at a given delta
|
||||
|
||||
Usage:
|
||||
python prerender.py eval
|
||||
python prerender.py ou --rho 0.95
|
||||
python prerender.py traj --delta 16 --h5_path data/reacher.h5
|
||||
|
||||
Saves images as uint8 (3, 64, 64) to keep disk usage ~1.2 GB per 100k images.
|
||||
Normalization stats computed and saved; applied at training time.
|
||||
"""
|
||||
|
||||
import os
|
||||
os.environ.setdefault("MUJOCO_GL", "egl")
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from scipy.stats import pearsonr, shapiro, skew, kurtosis
|
||||
from tqdm import tqdm
|
||||
from dm_control import suite
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# RENDERING
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
TARGET = np.array([0.1, 0.1])
|
||||
IMG_SIZE = 64
|
||||
|
||||
|
||||
def make_env():
|
||||
return suite.load(domain_name="reacher", task_name="hard")
|
||||
|
||||
|
||||
def render_at(env, qpos, height=IMG_SIZE, width=IMG_SIZE):
|
||||
"""Render → (3, H, W) uint8."""
|
||||
env.physics.data.qpos[:2] = qpos
|
||||
env.physics.data.qvel[:] = 0
|
||||
env.physics.named.model.geom_pos['target', :2] = TARGET
|
||||
env.physics.forward()
|
||||
rgb = env.physics.render(height=height, width=width, camera_id=0)
|
||||
return rgb.transpose(2, 0, 1) # uint8, (3, H, W)
|
||||
|
||||
|
||||
def render_batch(env, qpos_batch):
|
||||
"""Render → (N, 3, H, W) uint8."""
|
||||
N = len(qpos_batch)
|
||||
imgs = np.empty((N, 3, IMG_SIZE, IMG_SIZE), dtype=np.uint8)
|
||||
for i in tqdm(range(N), desc="Rendering"):
|
||||
imgs[i] = render_at(env, qpos_batch[i])
|
||||
return imgs
|
||||
|
||||
|
||||
def compute_norm_stats(imgs_uint8):
|
||||
"""Compute per-channel mean/std from uint8 images. Returns float32 arrays."""
|
||||
imgs = imgs_uint8.astype(np.float32) / 255.0
|
||||
mean = imgs.mean(axis=(0, 2, 3)) # (3,)
|
||||
std = imgs.std(axis=(0, 2, 3)) # (3,)
|
||||
return mean.astype(np.float32), std.astype(np.float32)
|
||||
|
||||
|
||||
def save_dataset(out_dir, z_t, z_tp1, img_t, img_tp1, meta):
|
||||
"""Save arrays + metadata to directory."""
|
||||
out_dir = Path(out_dir)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
np.save(out_dir / "z_t.npy", z_t)
|
||||
np.save(out_dir / "z_tp1.npy", z_tp1)
|
||||
np.save(out_dir / "img_t.npy", img_t)
|
||||
np.save(out_dir / "img_tp1.npy", img_tp1)
|
||||
|
||||
# Norm stats from img_t
|
||||
mean, std = compute_norm_stats(img_t)
|
||||
np.save(out_dir / "img_mean.npy", mean)
|
||||
np.save(out_dir / "img_std.npy", std)
|
||||
|
||||
meta["img_mean"] = mean.tolist()
|
||||
meta["img_std"] = std.tolist()
|
||||
with open(out_dir / "meta.json", "w") as f:
|
||||
json.dump(meta, f, indent=2)
|
||||
|
||||
size_gb = sum(
|
||||
(out_dir / fn).stat().st_size for fn in
|
||||
["img_t.npy", "img_tp1.npy", "z_t.npy", "z_tp1.npy"]
|
||||
) / 1e9
|
||||
print(f" Saved to {out_dir} ({size_gb:.2f} GB)")
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# EVAL
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def prerender_eval(args):
|
||||
"""10k i.i.d. Gaussian samples + rendered images."""
|
||||
out_dir = Path(args.data_root) / "eval"
|
||||
if (out_dir / "img.npy").exists() and not args.force:
|
||||
print(f"Eval data already exists at {out_dir}, skipping (use --force)")
|
||||
return
|
||||
|
||||
rng = np.random.default_rng(args.eval_seed)
|
||||
z = rng.standard_normal((args.n_eval, 2)).astype(np.float32)
|
||||
|
||||
env = make_env()
|
||||
print(f"Rendering {args.n_eval} eval images...")
|
||||
imgs = render_batch(env, z)
|
||||
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
np.save(out_dir / "z.npy", z)
|
||||
np.save(out_dir / "img.npy", imgs)
|
||||
|
||||
mean, std = compute_norm_stats(imgs)
|
||||
np.save(out_dir / "img_mean.npy", mean)
|
||||
np.save(out_dir / "img_std.npy", std)
|
||||
|
||||
with open(out_dir / "meta.json", "w") as f:
|
||||
json.dump({"n_eval": args.n_eval, "seed": args.eval_seed,
|
||||
"img_mean": mean.tolist(), "img_std": std.tolist()}, f, indent=2)
|
||||
print(f" Saved to {out_dir}")
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# OU
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def prerender_ou(args):
|
||||
"""100k OU pairs for a given rho."""
|
||||
rho = args.rho
|
||||
out_dir = Path(args.data_root) / "ou" / f"rho={rho:.2f}"
|
||||
if (out_dir / "img_t.npy").exists() and not args.force:
|
||||
print(f"OU data for rho={rho} already exists, skipping (use --force)")
|
||||
return
|
||||
|
||||
N = args.n_train
|
||||
rng = np.random.default_rng(args.render_seed)
|
||||
z_t = rng.standard_normal((N, 2)).astype(np.float32)
|
||||
eps = rng.standard_normal((N, 2)).astype(np.float32)
|
||||
z_tp1 = rho * z_t + np.sqrt(1 - rho**2) * eps
|
||||
|
||||
env = make_env()
|
||||
print(f"OU rho={rho}: rendering {2 * N} images...")
|
||||
img_t = render_batch(env, z_t)
|
||||
img_tp1 = render_batch(env, z_tp1)
|
||||
|
||||
meta = {"type": "ou", "rho": rho, "n": N, "seed": args.render_seed}
|
||||
save_dataset(out_dir, z_t, z_tp1, img_t, img_tp1, meta)
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# TRAJECTORY
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def load_episodes(h5_path):
|
||||
"""Load qpos grouped by episode → (n_episodes, T, 2)."""
|
||||
import h5py
|
||||
with h5py.File(h5_path, "r") as f:
|
||||
qpos = np.array(f["qpos"])
|
||||
ep_len = np.array(f["ep_len"])
|
||||
T = ep_len[0]
|
||||
assert (ep_len == T).all(), f"Non-uniform episode lengths"
|
||||
episodes = qpos.reshape(len(ep_len), T, 2)
|
||||
print(f"Loaded {len(episodes)} episodes, {T} steps each")
|
||||
return episodes
|
||||
|
||||
|
||||
def subsample_pairs(episodes, delta, n_per_episode, seed):
|
||||
"""Sample n_per_episode (t, t+delta) pairs from each episode."""
|
||||
rng = np.random.default_rng(seed)
|
||||
n_ep, T, d = episodes.shape
|
||||
max_start = T - delta
|
||||
z_t_list, z_tp1_list = [], []
|
||||
for ep in episodes:
|
||||
starts = rng.choice(max_start, size=n_per_episode, replace=False)
|
||||
z_t_list.append(ep[starts])
|
||||
z_tp1_list.append(ep[starts + delta])
|
||||
return (np.concatenate(z_t_list).astype(np.float32),
|
||||
np.concatenate(z_tp1_list).astype(np.float32))
|
||||
|
||||
|
||||
def traj_diagnostics(episodes, delta):
|
||||
"""Compute autocorrelation + normality stats."""
|
||||
n_ep, T, d = episodes.shape
|
||||
ms = T - delta
|
||||
z_t = episodes[:, :ms].reshape(-1, d)
|
||||
z_tp1 = episodes[:, delta:delta+ms].reshape(-1, d)
|
||||
|
||||
diag = {"delta": delta}
|
||||
for i, name in enumerate(["shoulder", "wrist"]):
|
||||
r, _ = pearsonr(z_t[:, i], z_tp1[:, i])
|
||||
diag[f"rho_{name}"] = float(r)
|
||||
diag[f"skew_{name}"] = float(skew(z_t[:, i]))
|
||||
diag[f"kurtosis_{name}"] = float(kurtosis(z_t[:, i]))
|
||||
sub = z_t[np.random.choice(len(z_t), 5000, replace=False), i]
|
||||
_, p = shapiro(sub)
|
||||
diag[f"shapiro_p_{name}"] = float(p)
|
||||
diag["rho_mean"] = (diag["rho_shoulder"] + diag["rho_wrist"]) / 2
|
||||
return diag
|
||||
|
||||
|
||||
def prerender_traj(args):
|
||||
"""100k pairs subsampled from LeWM trajectories at a given delta."""
|
||||
delta = args.delta
|
||||
out_dir = Path(args.data_root) / "traj" / f"delta={delta}"
|
||||
if (out_dir / "img_t.npy").exists() and not args.force:
|
||||
print(f"Traj data for delta={delta} already exists, skipping")
|
||||
return
|
||||
|
||||
episodes = load_episodes(args.h5_path)
|
||||
n_episodes = len(episodes)
|
||||
n_per_episode = args.n_train // n_episodes
|
||||
N_actual = n_per_episode * n_episodes
|
||||
print(f"delta={delta}: {n_per_episode} pairs/episode × {n_episodes} = {N_actual}")
|
||||
|
||||
# Diagnostics
|
||||
diag = traj_diagnostics(episodes, delta)
|
||||
print(f" rho: shoulder={diag['rho_shoulder']:.4f}, "
|
||||
f"wrist={diag['rho_wrist']:.4f}")
|
||||
print(f" skew: {diag['skew_shoulder']:.3f}, {diag['skew_wrist']:.3f}")
|
||||
|
||||
# Subsample
|
||||
z_t, z_tp1 = subsample_pairs(episodes, delta, n_per_episode, args.render_seed)
|
||||
|
||||
# Render
|
||||
env = make_env()
|
||||
print(f" Rendering {2 * len(z_t)} images...")
|
||||
img_t = render_batch(env, z_t)
|
||||
img_tp1 = render_batch(env, z_tp1)
|
||||
|
||||
meta = {"type": "traj", "delta": delta, "n": len(z_t),
|
||||
"n_per_episode": n_per_episode, "seed": args.render_seed,
|
||||
**diag}
|
||||
save_dataset(out_dir, z_t, z_tp1, img_t, img_tp1, meta)
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# MAIN
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
sub = p.add_subparsers(dest="mode", required=True)
|
||||
|
||||
# Shared
|
||||
for name in ["eval", "ou", "traj"]:
|
||||
sp = sub.add_parser(name)
|
||||
sp.add_argument("--data_root", type=str, default="data/reacher")
|
||||
sp.add_argument("--force", action="store_true")
|
||||
sp.add_argument("--render_seed", type=int, default=9999)
|
||||
|
||||
# eval
|
||||
sub.choices["eval"].add_argument("--n_eval", type=int, default=10000)
|
||||
sub.choices["eval"].add_argument("--eval_seed", type=int, default=8888)
|
||||
|
||||
# ou
|
||||
sub.choices["ou"].add_argument("--rho", type=float, required=True)
|
||||
sub.choices["ou"].add_argument("--n_train", type=int, default=100000)
|
||||
|
||||
# traj
|
||||
sub.choices["traj"].add_argument("--delta", type=int, required=True)
|
||||
sub.choices["traj"].add_argument("--h5_path", type=str, required=True)
|
||||
sub.choices["traj"].add_argument("--n_train", type=int, default=100000)
|
||||
|
||||
args = p.parse_args()
|
||||
|
||||
if args.mode == "eval":
|
||||
prerender_eval(args)
|
||||
elif args.mode == "ou":
|
||||
prerender_ou(args)
|
||||
elif args.mode == "traj":
|
||||
prerender_traj(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,324 @@
|
||||
"""
|
||||
Unified experiment runner. Loads config YAML, builds mixing + encoder,
|
||||
calls engine.train_and_evaluate, saves standardized .pt output.
|
||||
|
||||
Usage:
|
||||
python run.py --config configs/2d.yaml --run spiral --seed 1337
|
||||
python run.py --config configs/ablation.yaml --run spiral_lejepa --seed 1337
|
||||
python run.py --config configs/scaling.yaml --N 16 --seed 0
|
||||
python run.py --config configs/grid.yaml --lamb 0.01 --rho 0.9 --seed 0
|
||||
"""
|
||||
|
||||
import argparse, os, json, yaml
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from lejepa_id.mixing import MIXINGS_2D, make_coupling_mixing
|
||||
from lejepa_id.models import make_mlp_encoder, make_matched_encoder
|
||||
from lejepa_id.data import sample_latents, ou_augment
|
||||
from lejepa_id.metrics import compute_all_metrics, compute_recovery_metrics
|
||||
from lejepa_id.engine import train_and_evaluate
|
||||
|
||||
|
||||
def _jsonify(obj):
|
||||
"""Convert numpy types to Python natives for JSON serialization."""
|
||||
if isinstance(obj, dict):
|
||||
return {k: _jsonify(v) for k, v in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [_jsonify(v) for v in obj]
|
||||
elif isinstance(obj, (np.floating,)):
|
||||
return float(obj)
|
||||
elif isinstance(obj, (np.integer,)):
|
||||
return int(obj)
|
||||
elif isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
return obj
|
||||
|
||||
def build_mixing(mixing_name, N, n_layers=4, seed=1337, device="cuda"):
|
||||
"""Build mixing function from name."""
|
||||
if mixing_name in MIXINGS_2D:
|
||||
return MIXINGS_2D[mixing_name]
|
||||
elif mixing_name in ("nvp", "coupling"):
|
||||
return make_coupling_mixing(N, n_layers=n_layers, seed=seed, device=device)
|
||||
else:
|
||||
raise ValueError(f"Unknown mixing: {mixing_name}")
|
||||
|
||||
|
||||
def build_encoder(encoder_type, N, hidden=256, n_layers=4, seed=42, device="cuda"):
|
||||
"""Build encoder from type string."""
|
||||
if encoder_type == "mlp":
|
||||
return make_mlp_encoder(N, hidden=hidden, device=device)
|
||||
elif encoder_type == "matched":
|
||||
return make_matched_encoder(N, n_layers=n_layers, seed=seed, device=device)
|
||||
else:
|
||||
raise ValueError(f"Unknown encoder: {encoder_type}")
|
||||
|
||||
|
||||
def resolve_run_spec(cfg, args):
|
||||
"""Resolve the full run specification from config + CLI args.
|
||||
Returns a dict with all parameters needed for one training run."""
|
||||
experiment = cfg["experiment"]
|
||||
|
||||
# Start with config-level defaults
|
||||
spec = {
|
||||
"experiment": experiment,
|
||||
"N": cfg.get("N", 2),
|
||||
"source_dist": cfg.get("source_dist", "gaussian"),
|
||||
"num_eval": cfg.get("num_eval", 10000),
|
||||
"steps": cfg.get("steps", 10000),
|
||||
"lr": cfg.get("lr", 3e-3),
|
||||
"batch_size": cfg.get("batch_size", 256),
|
||||
"rho": cfg.get("rho", 0.95),
|
||||
"lamb": cfg.get("lamb"),
|
||||
"sigma": cfg.get("sigma", 1.0),
|
||||
"source_alpha": cfg.get("source_alpha"), # NEW
|
||||
"log_every": cfg.get("log_every", 100),
|
||||
"encoder": cfg.get("encoder", "mlp"),
|
||||
"hidden": cfg.get("hidden", 256),
|
||||
"n_layers": cfg.get("n_layers", 4),
|
||||
"mixing": cfg.get("mixing", "spiral"),
|
||||
"mode": cfg.get("mode", "lejepa"),
|
||||
"seed": args.seed,
|
||||
}
|
||||
|
||||
if experiment in ("2d", "ablation"):
|
||||
# Look up run-specific overrides
|
||||
run_name = args.run
|
||||
run_cfg = cfg["runs"][run_name]
|
||||
spec["run_name"] = run_name
|
||||
for k in ("mixing", "encoder", "hidden", "n_layers", "mode", "lamb", "sigma"):
|
||||
if k in run_cfg:
|
||||
spec[k] = run_cfg[k]
|
||||
|
||||
elif experiment == "scaling":
|
||||
N = args.N
|
||||
spec["N"] = N
|
||||
spec["mixing"] = "coupling"
|
||||
if args.mode is not None:
|
||||
spec["mode"] = args.mode
|
||||
# Mode-specific lamb (whiten uses different default)
|
||||
if spec["mode"] == "whiten":
|
||||
spec["lamb"] = cfg.get("lamb_whiten", 0.5)
|
||||
spec["run_name"] = f"N={N}_{spec['mode']}"
|
||||
|
||||
elif experiment == "grid":
|
||||
spec["lamb"] = args.lamb
|
||||
spec["rho"] = args.rho
|
||||
spec["run_name"] = f"lamb={args.lamb:.0e}_rho={args.rho:.2f}"
|
||||
|
||||
elif experiment == "gennorm":
|
||||
if args.alpha is None:
|
||||
raise ValueError("--alpha required for gennorm experiment")
|
||||
spec["source_dist"] = "gennorm"
|
||||
spec["source_alpha"] = args.alpha
|
||||
run_name = args.run
|
||||
run_cfg = cfg["runs"][run_name]
|
||||
for k in ("mixing", "encoder", "hidden", "n_layers", "mode", "lamb", "sigma"):
|
||||
if k in run_cfg:
|
||||
spec[k] = run_cfg[k]
|
||||
spec["run_name"] = f"{run_name}_alpha={args.alpha:g}"
|
||||
|
||||
return spec
|
||||
|
||||
|
||||
def run_single(spec, device):
|
||||
"""Execute one training run from a resolved spec. Returns result dict."""
|
||||
N = spec["N"]
|
||||
seed = spec["seed"]
|
||||
|
||||
torch.manual_seed(seed)
|
||||
np.random.seed(seed)
|
||||
|
||||
# Build mixing
|
||||
mix_seed = seed
|
||||
n_layers = spec.get("n_layers", 4)
|
||||
mix_fn = build_mixing(spec["mixing"], N, n_layers=n_layers,
|
||||
seed=mix_seed, device=device)
|
||||
|
||||
# Build encoder (different seed from mixing)
|
||||
enc_seed = seed + 77777
|
||||
encoder = build_encoder(spec["encoder"], N, hidden=spec.get("hidden", 256),
|
||||
n_layers=n_layers, seed=enc_seed, device=device)
|
||||
|
||||
# Fixed eval set
|
||||
z_eval = sample_latents(spec["num_eval"], N, dist=spec["source_dist"],
|
||||
device=device, alpha=spec.get("source_alpha"))
|
||||
|
||||
# Train
|
||||
encoder, log = train_and_evaluate(
|
||||
encoder, mix_fn,
|
||||
N=N, rho=spec["rho"], lamb=spec["lamb"], mode=spec["mode"],
|
||||
source_dist=spec["source_dist"],
|
||||
source_alpha=spec.get("source_alpha"),
|
||||
sigma=spec["sigma"],
|
||||
steps=spec["steps"], batch_size=spec["batch_size"], lr=spec["lr"],
|
||||
z_eval=z_eval, log_every=spec["log_every"], device=device,
|
||||
)
|
||||
|
||||
# Final metrics from 10k eval set
|
||||
encoder.eval()
|
||||
with torch.no_grad():
|
||||
x_eval = mix_fn(z_eval)
|
||||
h_eval = encoder(x_eval)
|
||||
# z_prime = ou_augment(z_eval, spec["rho"], n_views=1).squeeze(0)
|
||||
z_prime = ou_augment(
|
||||
z_eval, spec["rho"], n_views=1,
|
||||
dist=spec["source_dist"],
|
||||
alpha=spec.get("source_alpha")
|
||||
).squeeze(0)
|
||||
h_prime = encoder(mix_fn(z_prime))
|
||||
|
||||
final_metrics = compute_all_metrics(
|
||||
z_eval, mix_fn(z_eval), h_eval, h_prime, spec["rho"], N,
|
||||
)
|
||||
|
||||
# Fixed-grid evaluation (cross-distribution comparable, only for 2D)
|
||||
if N == 2:
|
||||
with torch.no_grad():
|
||||
g = torch.linspace(-3.0, 3.0, 100, device=device)
|
||||
z_grid = torch.stack(torch.meshgrid(g, g, indexing='ij'), dim=-1).reshape(-1, N)
|
||||
h_grid = encoder(mix_fn(z_grid))
|
||||
final_metrics.update(compute_recovery_metrics(z_grid, h_grid, N, suffix="_grid"))
|
||||
|
||||
# Large scatter data for plotting (only for 2d/ablation)
|
||||
if spec["experiment"] in ("2d", "ablation"):
|
||||
with torch.no_grad():
|
||||
z_plot = sample_latents(100000, N, dist=spec["source_dist"],
|
||||
device=device, alpha=spec.get("source_alpha"))
|
||||
x_plot = mix_fn(z_plot)
|
||||
h_chunks = []
|
||||
for i in range(0, len(z_plot), 10000):
|
||||
h_chunks.append(encoder(x_plot[i:i+10000]))
|
||||
h_plot = torch.cat(h_chunks, dim=0)
|
||||
z_np = z_plot.cpu().numpy()
|
||||
x_np = x_plot.cpu().numpy()
|
||||
h_np = h_plot.cpu().numpy()
|
||||
else:
|
||||
z_np, x_np, h_np = None, None, None
|
||||
|
||||
# JSON-serializable result (scalars + training curves)
|
||||
result = {
|
||||
# Identity
|
||||
"experiment": spec["experiment"],
|
||||
"run_name": spec["run_name"],
|
||||
"mixing": spec["mixing"],
|
||||
"encoder": spec["encoder"],
|
||||
"mode": spec["mode"],
|
||||
"source_dist": spec["source_dist"],
|
||||
"source_alpha": spec.get("source_alpha"),
|
||||
"seed": seed,
|
||||
"N": N,
|
||||
# Hyperparameters
|
||||
"lamb": spec["lamb"],
|
||||
"rho": spec["rho"],
|
||||
"lr": spec["lr"],
|
||||
"steps": spec["steps"],
|
||||
"batch_size": spec["batch_size"],
|
||||
"n_layers": n_layers,
|
||||
"hidden": spec.get("hidden", None),
|
||||
# Final metrics
|
||||
**final_metrics,
|
||||
"final_align": log["align"][-1],
|
||||
"final_sigreg": log["sigreg"][-1],
|
||||
"final_whiten": log["whiten"][-1],
|
||||
"final_loss": log["total"][-1],
|
||||
# Training curves
|
||||
"log": log,
|
||||
}
|
||||
|
||||
# Heavy data (arrays + model) — only saved as .pt for 2d/ablation
|
||||
arrays = {
|
||||
"z": z_np, "x": x_np, "h": h_np,
|
||||
"model_state_dict": encoder.state_dict(),
|
||||
}
|
||||
|
||||
return result, arrays
|
||||
|
||||
|
||||
def save_result(result, arrays, out_dir, fname_base, save_pt=False):
|
||||
"""Save JSON always; .pt with arrays/model only when requested."""
|
||||
# JSON
|
||||
json_path = os.path.join(out_dir, fname_base + ".json")
|
||||
with open(json_path, "w") as f:
|
||||
json.dump(_jsonify(result), f, indent=2)
|
||||
print(f"Saved {fname_base}.json")
|
||||
|
||||
# .pt (arrays + model) for 2d/ablation scatter plots
|
||||
if save_pt:
|
||||
pt_path = os.path.join(out_dir, fname_base + ".pt")
|
||||
torch.save({**result, **arrays}, pt_path)
|
||||
print(f"Saved {fname_base}.pt")
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--config", type=str, required=True)
|
||||
# Sweep variables (CLI overrides)
|
||||
p.add_argument("--run", type=str, default=None, help="Run name (2d/ablation)")
|
||||
p.add_argument("--seed", type=int, required=True)
|
||||
p.add_argument("--N", type=int, default=None, help="Latent dim (scaling)")
|
||||
p.add_argument("--lamb", type=float, default=None, help="Lambda (grid)")
|
||||
p.add_argument("--rho", type=float, default=None, help="Rho (grid)")
|
||||
p.add_argument("--alpha", type=float, default=None, help="Gennorm shape (gennorm)")
|
||||
p.add_argument("--mode", type=str, default=None,
|
||||
help="Override mode (lejepa/whiten/infonce)")
|
||||
args = p.parse_args()
|
||||
|
||||
with open(args.config) as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f"Device: {device}")
|
||||
|
||||
spec = resolve_run_spec(cfg, args)
|
||||
out_dir = cfg["out"]
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
|
||||
experiment = cfg["experiment"]
|
||||
save_pt = experiment in ("2d", "ablation")
|
||||
|
||||
if experiment == "scaling":
|
||||
# For small N, train K encoders, pick best
|
||||
K = cfg.get("K", 1)
|
||||
# For large N, all converge
|
||||
if spec["N"] > 32:
|
||||
K = 1
|
||||
best_result = None
|
||||
best_arrays = None
|
||||
best_loss = float("inf")
|
||||
|
||||
for k in range(K):
|
||||
spec_k = dict(spec)
|
||||
spec_k["seed"] = spec["seed"] + k * 1000
|
||||
print(f"\n Encoder {k+1}/{K} (seed={spec_k['seed']})")
|
||||
result, arrays = run_single(spec_k, device)
|
||||
print(f" R²(h->z)={result['r2_hz']:.4f} "
|
||||
f"orth={result['orth_err']:.4f} "
|
||||
f"loss={result['final_loss']:.6f}")
|
||||
|
||||
if result["final_loss"] < best_loss:
|
||||
best_loss = result["final_loss"]
|
||||
best_result = result
|
||||
best_arrays = arrays
|
||||
|
||||
best_result["K"] = K
|
||||
best_result["seed"] = spec["seed"] # original seed
|
||||
fname = f"{spec['run_name']}_seed={spec['seed']}"
|
||||
save_result(best_result, best_arrays, out_dir, fname, save_pt=False)
|
||||
print(f" R²(h->z)={best_result['r2_hz']:.4f} orth={best_result['orth_err']:.4f}")
|
||||
|
||||
else:
|
||||
# Single run
|
||||
print(f"\n{'='*50}")
|
||||
print(f"{spec['run_name']} seed={spec['seed']}")
|
||||
print(f"{'='*50}")
|
||||
result, arrays = run_single(spec, device)
|
||||
|
||||
fname = f"{spec['run_name']}_seed={spec['seed']}"
|
||||
save_result(result, arrays, out_dir, fname, save_pt=save_pt)
|
||||
print(f" R²(z->h)={result['r2_zh']:.4f} R²(h->z)={result['r2_hz']:.4f} "
|
||||
f"orth={result['orth_err']:.4f}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,361 @@
|
||||
"""
|
||||
Train LeJEPA on prerendered Reacher data.
|
||||
|
||||
Loads images + latents from disk (output of prerender.py),
|
||||
sweeps lambda × seed. Final model is always used (no K inits, no loss selection).
|
||||
|
||||
Works identically for OU and trajectory data — just point --data_dir
|
||||
at the right directory.
|
||||
|
||||
Usage:
|
||||
python run_reacher.py --config configs/reacher.yaml \
|
||||
--data_dir data/reacher/ou/rho=0.95
|
||||
python run_reacher.py --config configs/reacher.yaml \
|
||||
--data_dir data/reacher/traj/delta=16
|
||||
"""
|
||||
|
||||
import argparse, os, json, yaml
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from lejepa_id.losses import SIGReg, alignment_loss
|
||||
from lejepa_id.models import make_cnn_encoder
|
||||
from lejepa_id.metrics import bidirectional_r2
|
||||
|
||||
from sklearn.linear_model import LinearRegression
|
||||
from scipy.linalg import orthogonal_procrustes
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# DATA LOADING
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def load_dataset(data_dir):
|
||||
"""Load prerendered (img, z) pairs. Images stored as uint8."""
|
||||
data_dir = str(data_dir)
|
||||
z_t = np.load(os.path.join(data_dir, "z_t.npy"))
|
||||
z_tp1 = np.load(os.path.join(data_dir, "z_tp1.npy"))
|
||||
img_t = np.load(os.path.join(data_dir, "img_t.npy"))
|
||||
img_tp1 = np.load(os.path.join(data_dir, "img_tp1.npy"))
|
||||
mean = np.load(os.path.join(data_dir, "img_mean.npy"))
|
||||
std = np.load(os.path.join(data_dir, "img_std.npy"))
|
||||
with open(os.path.join(data_dir, "meta.json")) as f:
|
||||
meta = json.load(f)
|
||||
print(f"Loaded {len(z_t)} pairs from {data_dir}")
|
||||
return z_t, z_tp1, img_t, img_tp1, mean, std, meta
|
||||
|
||||
|
||||
def load_eval(eval_dir):
|
||||
"""Load prerendered eval set."""
|
||||
eval_dir = str(eval_dir)
|
||||
z = np.load(os.path.join(eval_dir, "z.npy"))
|
||||
img = np.load(os.path.join(eval_dir, "img.npy"))
|
||||
mean = np.load(os.path.join(eval_dir, "img_mean.npy"))
|
||||
std = np.load(os.path.join(eval_dir, "img_std.npy"))
|
||||
print(f"Loaded {len(z)} eval samples")
|
||||
return z, img, mean, std
|
||||
|
||||
|
||||
def normalize_uint8(img_uint8, mean, std):
|
||||
"""Convert uint8 → float32 normalized. mean/std are (3,) arrays."""
|
||||
img = img_uint8.astype(np.float32) / 255.0
|
||||
img = (img - mean[None, :, None, None]) / (std[None, :, None, None] + 1e-6)
|
||||
return img
|
||||
|
||||
|
||||
class ImageDataset(torch.utils.data.Dataset):
|
||||
"""Normalized float32 image pairs + latents."""
|
||||
def __init__(self, img_t, img_tp1, z_t, z_tp1, mean, std):
|
||||
self.img_t = torch.from_numpy(normalize_uint8(img_t, mean, std))
|
||||
self.img_tp1 = torch.from_numpy(normalize_uint8(img_tp1, mean, std))
|
||||
self.z_t = torch.from_numpy(z_t)
|
||||
self.z_tp1 = torch.from_numpy(z_tp1)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.img_t)
|
||||
|
||||
def __getitem__(self, i):
|
||||
return self.img_t[i], self.img_tp1[i], self.z_t[i], self.z_tp1[i]
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# TRAINING
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_embeddings(encoder, images, device, batch_size=512):
|
||||
encoder.eval()
|
||||
embeds = []
|
||||
for i in range(0, len(images), batch_size):
|
||||
batch = images[i:i+batch_size].to(device)
|
||||
embeds.append(encoder(batch).cpu())
|
||||
return torch.cat(embeds)
|
||||
|
||||
|
||||
def train_one(encoder, loader, eval_data, lamb, cfg, device):
|
||||
"""Train one encoder. Returns final model (no selection)."""
|
||||
sigreg = SIGReg(n_slices=cfg["n_slices"]).to(device)
|
||||
opt = torch.optim.AdamW(encoder.parameters(), lr=cfg["lr"], weight_decay=1e-4)
|
||||
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=cfg["epochs"])
|
||||
|
||||
eval_imgs, eval_z = eval_data
|
||||
|
||||
log = {"align": [], "sigreg": [], "total": [], "z_std": [], "r2_hz": []}
|
||||
|
||||
for epoch in range(cfg["epochs"]):
|
||||
encoder.train()
|
||||
ep = {k: [] for k in ["align", "sigreg", "total", "z_std"]}
|
||||
|
||||
for img_t, img_tp1, _, _ in loader:
|
||||
img_t, img_tp1 = img_t.to(device), img_tp1.to(device)
|
||||
z_t = encoder(img_t)
|
||||
z_tp1 = encoder(img_tp1)
|
||||
|
||||
h = torch.stack([z_t, z_tp1], dim=0)
|
||||
L_align = alignment_loss(h)
|
||||
L_sig = sigreg(h)
|
||||
loss = lamb * L_sig + (1 - lamb) * L_align
|
||||
|
||||
opt.zero_grad()
|
||||
loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(encoder.parameters(), 1.0)
|
||||
opt.step()
|
||||
|
||||
ep["total"].append(loss.item())
|
||||
ep["align"].append(L_align.item())
|
||||
ep["sigreg"].append(L_sig.item())
|
||||
with torch.no_grad():
|
||||
ep["z_std"].append(z_t.std(0).mean().item())
|
||||
|
||||
scheduler.step()
|
||||
|
||||
# Quick eval (on eval subset, for logging only)
|
||||
encoder.eval()
|
||||
h_eval = extract_embeddings(encoder, eval_imgs, device)
|
||||
_, r2_hz = bidirectional_r2(eval_z, h_eval)
|
||||
|
||||
for k in ep:
|
||||
log[k].append(float(np.mean(ep[k])))
|
||||
log["r2_hz"].append(r2_hz)
|
||||
|
||||
if (epoch + 1) % 10 == 0 or epoch == 0:
|
||||
print(f" epoch {epoch+1:3d}/{cfg['epochs']} "
|
||||
f"align={log['align'][-1]:.5f} "
|
||||
f"sig={log['sigreg'][-1]:.1f} "
|
||||
f"z_std={log['z_std'][-1]:.3f} "
|
||||
f"R²={r2_hz:.4f}")
|
||||
|
||||
return log
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# EVALUATION
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def final_eval(encoder, train_imgs, train_z, eval_imgs, eval_z, device):
|
||||
"""
|
||||
Full eval with proper train/test split.
|
||||
Fit linear regression on train embeddings, score on eval embeddings.
|
||||
"""
|
||||
h_train = extract_embeddings(encoder, train_imgs, device).numpy()
|
||||
h_eval = extract_embeddings(encoder, eval_imgs, device).numpy()
|
||||
z_train = train_z.numpy() if isinstance(train_z, torch.Tensor) else train_z
|
||||
z_eval = eval_z.numpy() if isinstance(eval_z, torch.Tensor) else eval_z
|
||||
|
||||
# Overall R² (fit on train, score on test)
|
||||
reg_hz = LinearRegression().fit(h_train, z_train)
|
||||
r2_hz = reg_hz.score(h_eval, z_eval)
|
||||
|
||||
reg_zh = LinearRegression().fit(z_train, h_train)
|
||||
r2_zh = reg_zh.score(z_eval, h_eval)
|
||||
|
||||
# Per-dimension R² (fit on train, score on test)
|
||||
r2_hz_per = []
|
||||
for i in range(z_train.shape[1]):
|
||||
reg_i = LinearRegression().fit(h_train, z_train[:, i])
|
||||
r2_hz_per.append(reg_i.score(h_eval, z_eval[:, i]))
|
||||
|
||||
# Sin/cos diagnostic (fit on train, score on test)
|
||||
z_train_sc = np.column_stack([np.sin(z_train), np.cos(z_train)])
|
||||
z_eval_sc = np.column_stack([np.sin(z_eval), np.cos(z_eval)])
|
||||
reg_sc = LinearRegression().fit(h_train, z_train_sc)
|
||||
r2_sincos = reg_sc.score(h_eval, z_eval_sc)
|
||||
|
||||
# Per-component sin/cos R²
|
||||
sincos_names = ["sin_shoulder", "cos_shoulder", "sin_wrist", "cos_wrist"]
|
||||
r2_sincos_per = {}
|
||||
for i, name in enumerate(sincos_names):
|
||||
reg_i = LinearRegression().fit(h_train, z_train_sc[:, i])
|
||||
r2_sincos_per[name] = reg_i.score(h_eval, z_eval_sc[:, i])
|
||||
|
||||
# Orthogonality error
|
||||
d = min(z_eval.shape[1], h_eval.shape[1])
|
||||
Zt = (z_eval[:, :d] - z_eval[:, :d].mean(0)).copy()
|
||||
Zl = (h_eval[:, :d] - h_eval[:, :d].mean(0)).copy()
|
||||
for Z in [Zt, Zl]:
|
||||
cov = np.cov(Z, rowvar=False)
|
||||
evals, evecs = np.linalg.eigh(cov)
|
||||
evals = np.maximum(evals, 1e-8)
|
||||
W = evecs @ np.diag(1 / np.sqrt(evals)) @ evecs.T
|
||||
Z[:] = Z @ W
|
||||
R, _ = orthogonal_procrustes(Zl, Zt)
|
||||
orth_err = float(np.linalg.norm(Zl @ R - Zt) / np.linalg.norm(Zt))
|
||||
|
||||
return {
|
||||
"r2_zh": r2_zh,
|
||||
"r2_hz": r2_hz,
|
||||
"r2_hz_per_dim": r2_hz_per,
|
||||
"r2_sincos": r2_sincos,
|
||||
"r2_sincos_per": r2_sincos_per,
|
||||
"orth_error": orth_err,
|
||||
"linear_map_W": reg_hz.coef_.T,
|
||||
"linear_map_b": reg_hz.intercept_,
|
||||
}
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
# MAIN
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _jsonify(obj):
|
||||
if isinstance(obj, dict):
|
||||
return {k: _jsonify(v) for k, v in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [_jsonify(v) for v in obj]
|
||||
elif isinstance(obj, (np.floating,)):
|
||||
return float(obj)
|
||||
elif isinstance(obj, (np.integer,)):
|
||||
return int(obj)
|
||||
elif isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
return obj
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--config", type=str, required=True)
|
||||
p.add_argument("--data_dir", type=str, required=True,
|
||||
help="Path to prerendered dataset (ou/rho=X or traj/delta=X)")
|
||||
args = p.parse_args()
|
||||
|
||||
with open(args.config) as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f"Device: {device}")
|
||||
|
||||
out_dir = cfg["out"]
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
|
||||
# ── Load data ────────────────────────────────────────────────────────
|
||||
z_t, z_tp1, img_t, img_tp1, train_mean, train_std, data_meta = \
|
||||
load_dataset(args.data_dir)
|
||||
|
||||
eval_dir = os.path.join(cfg["data_root"], "eval")
|
||||
z_eval, img_eval_u8, eval_mean, eval_std = load_eval(eval_dir)
|
||||
|
||||
# Normalize with training stats
|
||||
dataset = ImageDataset(img_t, img_tp1, z_t, z_tp1, train_mean, train_std)
|
||||
eval_imgs = torch.from_numpy(
|
||||
normalize_uint8(img_eval_u8, train_mean, train_std))
|
||||
eval_z = torch.from_numpy(z_eval)
|
||||
|
||||
# Train embeddings for linreg fitting (subsample for speed)
|
||||
n_fit = min(10000, len(dataset))
|
||||
fit_imgs = dataset.img_t[:n_fit]
|
||||
fit_z = dataset.z_t[:n_fit]
|
||||
|
||||
# Fast eval subset for in-training monitoring
|
||||
n_fast = cfg.get("n_eval_fast", 2000)
|
||||
eval_data_fast = (eval_imgs[:n_fast], eval_z[:n_fast])
|
||||
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset, batch_size=cfg["batch_size"], shuffle=True,
|
||||
num_workers=4, pin_memory=True, drop_last=True)
|
||||
|
||||
# Dataset label for output paths
|
||||
data_label = os.path.basename(args.data_dir)
|
||||
|
||||
# ── Sweep lambda × seed ──────────────────────────────────────────────
|
||||
all_results = []
|
||||
|
||||
for lamb in cfg["lambs"]:
|
||||
for seed in cfg["seeds"]:
|
||||
run_name = f"{data_label}_lamb={lamb:.0e}_seed={seed}"
|
||||
print(f"\n{'='*60}")
|
||||
print(f" {run_name}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
torch.manual_seed(seed)
|
||||
np.random.seed(seed)
|
||||
|
||||
encoder = make_cnn_encoder(
|
||||
d_latent=cfg["d_latent"], device=device)
|
||||
|
||||
log = train_one(
|
||||
encoder, loader, eval_data_fast,
|
||||
lamb=lamb, cfg=cfg, device=device)
|
||||
|
||||
# Full eval with train/test split
|
||||
metrics = final_eval(encoder, fit_imgs, fit_z,
|
||||
eval_imgs, eval_z, device)
|
||||
|
||||
result = {
|
||||
"experiment": "reacher",
|
||||
"run_name": run_name,
|
||||
"data_dir": args.data_dir,
|
||||
"lamb": lamb,
|
||||
"seed": seed,
|
||||
"d_latent": cfg["d_latent"],
|
||||
# Data meta (exclude 'seed' key to avoid overwriting training seed)
|
||||
**{k: v for k, v in data_meta.items() if k != "seed"},
|
||||
"render_seed": data_meta.get("seed", None),
|
||||
# Metrics
|
||||
**{k: v for k, v in metrics.items()
|
||||
if not isinstance(v, np.ndarray)},
|
||||
"best_r2_during_training": max(log["r2_hz"]),
|
||||
"final_r2_during_training": log["r2_hz"][-1],
|
||||
"final_loss": log["total"][-1],
|
||||
"final_align": log["align"][-1],
|
||||
"final_sigreg": log["sigreg"][-1],
|
||||
"log": log,
|
||||
}
|
||||
all_results.append(result)
|
||||
|
||||
print(f" → R²(h→z)={metrics['r2_hz']:.4f} "
|
||||
f"orth_err={metrics['orth_error']:.4f} "
|
||||
f"R²(sincos)={metrics['r2_sincos']:.4f}")
|
||||
print(f" per-dim R²: {['%.4f' % r for r in metrics['r2_hz_per_dim']]}")
|
||||
print(f" sincos: {metrics['r2_sincos_per']}")
|
||||
|
||||
# Save checkpoint + result
|
||||
run_dir = os.path.join(out_dir, run_name)
|
||||
os.makedirs(run_dir, exist_ok=True)
|
||||
torch.save({
|
||||
"encoder_state_dict": encoder.state_dict(),
|
||||
"train_mean": train_mean,
|
||||
"train_std": train_std,
|
||||
"d_latent": cfg["d_latent"],
|
||||
}, os.path.join(run_dir, "checkpoint.pt"))
|
||||
|
||||
with open(os.path.join(run_dir, "result.json"), "w") as f:
|
||||
json.dump(_jsonify(result), f, indent=2)
|
||||
|
||||
# ── Summary ──────────────────────────────────────────────────────────
|
||||
summary = {r["run_name"]: {k: v for k, v in r.items() if k != "log"}
|
||||
for r in all_results}
|
||||
with open(os.path.join(out_dir, f"summary_{data_label}.json"), "w") as f:
|
||||
json.dump(_jsonify(summary), f, indent=2)
|
||||
|
||||
print(f"\n{'data':>12s} {'lamb':>8s} {'seed':>4s} "
|
||||
f"{'R²(h→z)':>8s} {'R²(sc)':>8s} {'orth_err':>8s}")
|
||||
print("-" * 56)
|
||||
for r in all_results:
|
||||
print(f"{data_label:>12s} {r['lamb']:8.1e} {r['seed']:4d} "
|
||||
f"{r['r2_hz']:8.4f} {r['r2_sincos']:8.4f} "
|
||||
f"{r['orth_error']:8.4f}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,27 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=lejepa_2d
|
||||
#SBATCH --output=logs/2d_%A_%a.out
|
||||
#SBATCH --error=logs/2d_%A_%a.err
|
||||
#SBATCH --partition=gpuq
|
||||
#SBATCH --qos=slow_nice
|
||||
#SBATCH --gres=gpu:v100:1
|
||||
#SBATCH --cpus-per-task=4
|
||||
#SBATCH --mem=16G
|
||||
#SBATCH --time=02:00:00
|
||||
#SBATCH --array=0-7 # 8 runs; seeds loop inside
|
||||
|
||||
RUNS=(spiral_lejepa spiral_whiten banana_lejepa banana_whiten \
|
||||
sinusoid_lejepa sinusoid_whiten nvp_lejepa nvp_whiten)
|
||||
SEEDS=(1337 1338 1339)
|
||||
|
||||
RUN=${RUNS[$SLURM_ARRAY_TASK_ID]}
|
||||
|
||||
eval "$(conda shell.bash hook)"
|
||||
conda activate pytorch
|
||||
|
||||
mkdir -p logs
|
||||
for SEED in "${SEEDS[@]}"; do
|
||||
echo "${RUN} seed=${SEED}"
|
||||
python run.py --config configs/2d.yaml \
|
||||
--run "${RUN}" --seed "${SEED}"
|
||||
done
|
||||
@@ -0,0 +1,32 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=lejepa_gennorm
|
||||
#SBATCH --output=logs/gennorm_%A_%a.out
|
||||
#SBATCH --error=logs/gennorm_%A_%a.err
|
||||
#SBATCH --partition=gpuq
|
||||
#SBATCH --qos=slow_nice
|
||||
#SBATCH --gres=gpu:v100:1
|
||||
#SBATCH --cpus-per-task=4
|
||||
#SBATCH --mem=16G
|
||||
#SBATCH --time=03:00:00
|
||||
#SBATCH --array=0-71 # 8 runs x 9 alphas; seeds loop inside
|
||||
|
||||
RUNS=(spiral_lejepa spiral_whiten banana_lejepa banana_whiten \
|
||||
sinusoid_lejepa sinusoid_whiten nvp_lejepa nvp_whiten)
|
||||
ALPHAS=(0.125 0.25 0.5 1.0 2.0 4.0 8.0 16.0 32.0)
|
||||
SEEDS=(1337 1338 1339)
|
||||
|
||||
N_ALPHAS=${#ALPHAS[@]}
|
||||
RUN_IDX=$(( SLURM_ARRAY_TASK_ID / N_ALPHAS ))
|
||||
ALPHA_IDX=$(( SLURM_ARRAY_TASK_ID % N_ALPHAS ))
|
||||
RUN=${RUNS[$RUN_IDX]}
|
||||
ALPHA=${ALPHAS[$ALPHA_IDX]}
|
||||
|
||||
eval "$(conda shell.bash hook)"
|
||||
conda activate pytorch
|
||||
|
||||
mkdir -p logs
|
||||
for SEED in "${SEEDS[@]}"; do
|
||||
echo "${RUN} alpha=${ALPHA} seed=${SEED}"
|
||||
python run.py --config configs/gennorm.yaml \
|
||||
--run "${RUN}" --alpha "${ALPHA}" --seed "${SEED}"
|
||||
done
|
||||
@@ -0,0 +1,29 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=lejepa_grid
|
||||
#SBATCH --output=logs/grid_%A_%a.out
|
||||
#SBATCH --error=logs/grid_%A_%a.err
|
||||
#SBATCH --partition=gpuq
|
||||
#SBATCH --qos=slow_nice
|
||||
#SBATCH --gres=gpu:v100:1
|
||||
#SBATCH --cpus-per-task=4
|
||||
#SBATCH --mem=16G
|
||||
#SBATCH --time=12:00:00
|
||||
#SBATCH --array=0-8 # 9 lambda values; rhos and seeds loop inside
|
||||
|
||||
LAMBS=(1e-6 1e-5 1e-4 1e-3 5e-3 1e-2 5e-2 1e-1 5e-1)
|
||||
RHOS=(0.3 0.5 0.7 0.8 0.9 0.95 0.99)
|
||||
SEEDS=(0 1 2)
|
||||
|
||||
LAMB=${LAMBS[$SLURM_ARRAY_TASK_ID]}
|
||||
|
||||
eval "$(conda shell.bash hook)"
|
||||
conda activate pytorch
|
||||
|
||||
mkdir -p logs
|
||||
for RHO in "${RHOS[@]}"; do
|
||||
for SEED in "${SEEDS[@]}"; do
|
||||
echo "lamb=${LAMB} rho=${RHO} seed=${SEED}"
|
||||
python run.py --config configs/grid.yaml \
|
||||
--lamb "${LAMB}" --rho "${RHO}" --seed "${SEED}"
|
||||
done
|
||||
done
|
||||
@@ -0,0 +1,35 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=lejepa_ou
|
||||
#SBATCH --output=logs/ou_%A_%a.out
|
||||
#SBATCH --error=logs/ou_%A_%a.err
|
||||
#SBATCH --partition=gpuq
|
||||
#SBATCH --qos=slow_nice
|
||||
#SBATCH --gres=gpu:v100:1
|
||||
#SBATCH --cpus-per-task=4
|
||||
#SBATCH --mem=32G
|
||||
#SBATCH --time=24:00:00
|
||||
#SBATCH --array=0-6
|
||||
|
||||
# Each task: prerender eval (skipped if exists) + 200k images for one rho,
|
||||
# then train 4 lambdas × 3 seeds × 3 inits = 36 training runs
|
||||
|
||||
RHOS=(0.3 0.5 0.7 0.8 0.9 0.95 0.99)
|
||||
RHO_RAW=${RHOS[$SLURM_ARRAY_TASK_ID]}
|
||||
RHO=$(printf "%.2f" $RHO_RAW)
|
||||
|
||||
eval "$(conda shell.bash hook)"
|
||||
conda activate pytorch
|
||||
export MUJOCO_GL=egl
|
||||
mkdir -p logs
|
||||
|
||||
echo "Node: $(hostname) | rho=${RHO} | Start: $(date)"
|
||||
|
||||
# Step 1: Prerender (eval + this rho)
|
||||
python prerender.py eval
|
||||
python prerender.py ou --rho "${RHO}"
|
||||
|
||||
# Step 2: Train
|
||||
python run_reacher.py --config configs/reacher.yaml \
|
||||
--data_dir "data/reacher/ou/rho=${RHO}"
|
||||
|
||||
echo "Done: $(date)"
|
||||
@@ -0,0 +1,36 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=lejepa_traj
|
||||
#SBATCH --output=logs/traj_%A_%a.out
|
||||
#SBATCH --error=logs/traj_%A_%a.err
|
||||
#SBATCH --partition=gpuq
|
||||
#SBATCH --qos=slow_nice
|
||||
#SBATCH --gres=gpu:v100:1
|
||||
#SBATCH --cpus-per-task=4
|
||||
#SBATCH --mem=32G
|
||||
#SBATCH --time=24:00:00
|
||||
#SBATCH --array=0-6
|
||||
|
||||
# Each task: prerender eval (skipped if exists) + 200k images for one delta,
|
||||
# then train 4 lambdas × 3 seeds × 3 inits = 36 training runs
|
||||
|
||||
DELTAS=(1 2 4 8 16 32 64)
|
||||
DELTA=${DELTAS[$SLURM_ARRAY_TASK_ID]}
|
||||
|
||||
H5_PATH="data/reacher.h5"
|
||||
|
||||
eval "$(conda shell.bash hook)"
|
||||
conda activate pytorch
|
||||
export MUJOCO_GL=egl
|
||||
mkdir -p logs
|
||||
|
||||
echo "Node: $(hostname) | delta=${DELTA} | Start: $(date)"
|
||||
|
||||
# Step 1: Prerender (eval + this delta)
|
||||
python prerender.py eval
|
||||
python prerender.py traj --delta "${DELTA}" --h5_path "${H5_PATH}"
|
||||
|
||||
# Step 2: Train
|
||||
python run_reacher.py --config configs/reacher.yaml \
|
||||
--data_dir "data/reacher/traj/delta=${DELTA}"
|
||||
|
||||
echo "Done: $(date)"
|
||||
@@ -0,0 +1,30 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=lejepa_scale
|
||||
#SBATCH --output=logs/scale_%A_%a.out
|
||||
#SBATCH --error=logs/scale_%A_%a.err
|
||||
#SBATCH --partition=gpuq
|
||||
#SBATCH --qos=slow_nice
|
||||
#SBATCH --gres=gpu:v100:1
|
||||
#SBATCH --cpus-per-task=4
|
||||
#SBATCH --mem=16G
|
||||
#SBATCH --time=12:00:00 # ← bumped from 6h: 3 modes × 5 seeds = 15 runs per N
|
||||
#SBATCH --array=0-9 # 10 dims
|
||||
|
||||
DIMS=(2 4 8 16 32 64 128 256 512 1024)
|
||||
SEEDS=(0 1 2 3 4)
|
||||
MODES=(lejepa whiten infonce)
|
||||
|
||||
N=${DIMS[$SLURM_ARRAY_TASK_ID]}
|
||||
|
||||
eval "$(conda shell.bash hook)"
|
||||
conda activate pytorch
|
||||
|
||||
mkdir -p logs
|
||||
|
||||
for MODE in "${MODES[@]}"; do
|
||||
for SEED in "${SEEDS[@]}"; do
|
||||
echo "N=${N} seed=${SEED} mode=${MODE}"
|
||||
python -u run.py --config configs/scaling.yaml \
|
||||
--N "${N}" --seed "${SEED}" --mode "${MODE}"
|
||||
done
|
||||
done
|
||||
Reference in New Issue
Block a user