From 39ea1acbb5ad62043cb60c0f96b3bcb109607831 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Fri, 5 Jun 2026 16:03:17 -0500 Subject: Test first order FA tangent hierarchy --- scripts/fa_tangent_hierarchy_derivative_probe.py | 278 ++++++++++++++++++++++ scripts/plot_tangent_hierarchy_first_order_try.py | 166 +++++++++++++ 2 files changed, 444 insertions(+) create mode 100644 scripts/fa_tangent_hierarchy_derivative_probe.py create mode 100644 scripts/plot_tangent_hierarchy_first_order_try.py (limited to 'scripts') diff --git a/scripts/fa_tangent_hierarchy_derivative_probe.py b/scripts/fa_tangent_hierarchy_derivative_probe.py new file mode 100644 index 0000000..010b370 --- /dev/null +++ b/scripts/fa_tangent_hierarchy_derivative_probe.py @@ -0,0 +1,278 @@ +#!/usr/bin/env python3 +"""Probe whether the true initial operator derivative predicts finite-time gaps.""" + +from __future__ import annotations + +import argparse +import csv +import sys +from dataclasses import asdict, dataclass +from pathlib import Path + +import numpy as np +import torch + +SCRIPT_DIR = Path(__file__).resolve().parent +if str(SCRIPT_DIR) not in sys.path: + sys.path.insert(0, str(SCRIPT_DIR)) + +import downstream_capacity_sweep as dcs # noqa: E402 +from compressed_operator_predictor import fixed_rollout, kernel_pair, loss_from_residual # noqa: E402 + + +@dataclass(frozen=True) +class DerivativeRow: + init_seed: int + feedback_seed: int + epsilon: float + empirical_gap: float + fixed_gap: float + derivative_gap: float + fixed_error: float + derivative_error: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--input-dim", type=int, default=16) + parser.add_argument("--output-dim", type=int, default=4) + parser.add_argument("--width", type=int, default=64) + parser.add_argument("--hidden-layers", type=int, default=2) + parser.add_argument("--train-samples", type=int, default=128) + parser.add_argument("--test-samples", type=int, default=512) + parser.add_argument("--target-steps", type=int, default=50) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--epsilon", type=float, default=0.05) + parser.add_argument("--init-seeds", type=int, default=3) + parser.add_argument("--feedback-seeds", type=int, default=6) + parser.add_argument("--data-seed", type=int, default=123) + parser.add_argument("--feedback-scale", choices=["relu", "fan-in", "unit"], default="relu") + parser.add_argument("--device", choices=["cpu", "cuda"], default="cpu") + parser.add_argument("--torch-threads", type=int, default=8) + parser.add_argument( + "--outdir", + type=Path, + default=Path("outputs/fa_tangent_hierarchy_derivative_probe"), + ) + return parser.parse_args() + + +def make_config(args: argparse.Namespace) -> dcs.RunConfig: + return dcs.RunConfig( + task="random", + input_dim=args.input_dim, + output_dim=args.output_dim, + teacher_rank=4, + teacher_width=64, + teacher_hidden_layers=2, + normalize_targets=False, + widths=[args.width], + train_samples=args.train_samples, + test_samples=args.test_samples, + probe_samples=8, + steps=args.target_steps, + lr=args.lr, + optimizer="sgd", + init_seeds=args.init_seeds, + feedback_seeds=args.feedback_seeds, + init_seed_offset=0, + feedback_seed_offset=0, + data_seed=args.data_seed, + noise_std=0.0, + feedback_scale=args.feedback_scale, + capacity_q=0.01, + jacobian_lambda_rel=1e-3, + skip_jacobian=True, + device=args.device, + torch_threads=args.torch_threads, + outdir=str(args.outdir), + plot=False, + ) + + +def initialize_weights(config: dcs.RunConfig, width: int, hidden_layers: int, seed: int) -> list[torch.Tensor]: + generator = torch.Generator(device=config.device) + generator.manual_seed(seed) + dims = [config.input_dim, *([width] * hidden_layers), config.output_dim] + weights: list[torch.Tensor] = [] + for layer, (fan_in, fan_out) in enumerate(zip(dims[:-1], dims[1:])): + scale = np.sqrt(2.0 / fan_in) if layer < len(dims) - 2 else 1.0 / np.sqrt(fan_in) + weights.append( + torch.randn( + fan_out, + fan_in, + generator=generator, + dtype=torch.float64, + device=config.device, + ) + * scale + ) + return weights + + +def init_feedback(config: dcs.RunConfig, width: int, hidden_layers: int, seed: int) -> list[torch.Tensor]: + generator = torch.Generator(device=config.device) + generator.manual_seed(seed) + shapes = [(width, width)] * max(hidden_layers - 1, 0) + [(width, config.output_dim)] + return [ + torch.randn( + rows, + cols, + generator=generator, + dtype=torch.float64, + device=config.device, + ) + * dcs.feedback_scale(rows, config.feedback_scale) + for rows, cols in shapes + ] + + +def train_to_steps( + weights: list[torch.Tensor], + x: torch.Tensor, + y: torch.Tensor, + lr: float, + steps: int, + feedback: list[torch.Tensor] | None, +) -> list[torch.Tensor]: + current = dcs.clone_weights(weights) + for _ in range(steps): + grads = dcs.gradients(current, x, y, feedback) + current = [weight - lr * grad for weight, grad in zip(current, grads)] + return current + + +def fractional_step( + weights: list[torch.Tensor], + x: torch.Tensor, + y: torch.Tensor, + lr: float, + epsilon: float, + feedback: list[torch.Tensor] | None, +) -> list[torch.Tensor]: + grads = dcs.gradients(weights, x, y, feedback) + return [weight - epsilon * lr * grad for weight, grad in zip(weights, grads)] + + +def linear_velocity_rollout( + kernel0: np.ndarray, + velocity: np.ndarray, + residual: np.ndarray, + lr: float, + steps: int, + samples: int, +) -> np.ndarray: + current = residual.copy() + scale = lr / samples + for t in range(steps): + kernel = kernel0 + t * velocity + current = current - scale * (kernel @ current) + return current + + +def run_one( + config: dcs.RunConfig, + hidden_layers: int, + x: torch.Tensor, + y: torch.Tensor, + init_seed: int, + feedback_seed: int, + epsilon: float, +) -> DerivativeRow: + width = config.widths[0] + initial = initialize_weights(config, width, hidden_layers, init_seed) + feedback = init_feedback(config, width, hidden_layers, feedback_seed) + with torch.no_grad(): + r0 = (dcs.predict(initial, x) - y).reshape(-1).cpu().numpy() + + kbp0, kfa0 = kernel_pair(initial, x, feedback) + fixed_bp = loss_from_residual( + fixed_rollout(kbp0, r0, config.lr, config.steps, config.train_samples), + config.train_samples, + ) + fixed_fa = loss_from_residual( + fixed_rollout(kfa0, r0, config.lr, config.steps, config.train_samples), + config.train_samples, + ) + + bp_eps = fractional_step(initial, x, y, config.lr, epsilon, feedback=None) + fa_eps = fractional_step(initial, x, y, config.lr, epsilon, feedback=feedback) + kbp_eps, _ = kernel_pair(bp_eps, x, feedback) + _, kfa_eps = kernel_pair(fa_eps, x, feedback) + v_bp = (kbp_eps - kbp0) / epsilon + v_fa = (kfa_eps - kfa0) / epsilon + derivative_bp = loss_from_residual( + linear_velocity_rollout(kbp0, v_bp, r0, config.lr, config.steps, config.train_samples), + config.train_samples, + ) + derivative_fa = loss_from_residual( + linear_velocity_rollout(kfa0, v_fa, r0, config.lr, config.steps, config.train_samples), + config.train_samples, + ) + + bp_target = train_to_steps(initial, x, y, config.lr, config.steps, feedback=None) + fa_target = train_to_steps(initial, x, y, config.lr, config.steps, feedback=feedback) + empirical_gap = dcs.mse(fa_target, x, y) - dcs.mse(bp_target, x, y) + fixed_gap = fixed_fa - fixed_bp + derivative_gap = derivative_fa - derivative_bp + return DerivativeRow( + init_seed=init_seed, + feedback_seed=feedback_seed, + epsilon=epsilon, + empirical_gap=empirical_gap, + fixed_gap=fixed_gap, + derivative_gap=derivative_gap, + fixed_error=fixed_gap - empirical_gap, + derivative_error=derivative_gap - empirical_gap, + ) + + +def write_rows(path: Path, rows: list[DerivativeRow]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(DerivativeRow.__annotations__.keys())) + writer.writeheader() + for row in rows: + writer.writerow(asdict(row)) + + +def main() -> None: + args = parse_args() + torch.set_num_threads(args.torch_threads) + config = make_config(args) + x_train, y_train, *_ = dcs.make_data(config) + rows: list[DerivativeRow] = [] + total = args.init_seeds * args.feedback_seeds + count = 0 + for init_index in range(args.init_seeds): + init_seed = 10_000 + init_index + for feedback_index in range(args.feedback_seeds): + feedback_seed = 100_000 + init_index * 1000 + feedback_index + count += 1 + print(f"[{count}/{total}] init={init_seed} feedback={feedback_seed}", flush=True) + rows.append( + run_one( + config, + args.hidden_layers, + x_train, + y_train, + init_seed, + feedback_seed, + args.epsilon, + ) + ) + args.outdir.mkdir(parents=True, exist_ok=True) + csv_path = args.outdir / "derivative_probe_rows.csv" + write_rows(csv_path, rows) + fixed = np.array([row.fixed_error for row in rows], dtype=np.float64) + derivative = np.array([row.derivative_error for row in rows], dtype=np.float64) + print(f"rows: {csv_path}") + print(f"fixed MAE={np.abs(fixed).mean():.6g}, bias={fixed.mean():.6g}") + print( + f"derivative MAE={np.abs(derivative).mean():.6g}, " + f"bias={derivative.mean():.6g}" + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/plot_tangent_hierarchy_first_order_try.py b/scripts/plot_tangent_hierarchy_first_order_try.py new file mode 100644 index 0000000..c74c8d2 --- /dev/null +++ b/scripts/plot_tangent_hierarchy_first_order_try.py @@ -0,0 +1,166 @@ +#!/usr/bin/env python3 +"""Plot first-order FA tangent-hierarchy predictor diagnostics.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument( + "--csv", + type=Path, + default=Path("outputs/fa_tangent_hierarchy_first_order_try/compressed_operator_rows.csv"), + ) + parser.add_argument( + "--outdir", + type=Path, + default=Path("outputs/fa_tangent_hierarchy_first_order_try"), + ) + return parser.parse_args() + + +def metrics(df: pd.DataFrame) -> pd.DataFrame: + rows = [] + for early, group in df.groupby("early_steps"): + for label, column in [ + ("fixed K0", "fixed_gap"), + ("linear velocity", "linear_gap"), + ("retangent", "retangent_gap"), + ("compressed", "compressed_gap"), + ]: + error = group[column] - group["empirical_gap"] + rows.append( + { + "early_steps": int(early), + "predictor": label, + "mae": float(error.abs().mean()), + "bias": float(error.mean()), + "rmse": float(np.sqrt(np.mean(np.square(error)))), + "corr": float(group[column].corr(group["empirical_gap"])), + } + ) + return pd.DataFrame(rows) + + +def plot_scatter(df: pd.DataFrame, outdir: Path) -> Path: + early_values = sorted(df["early_steps"].unique()) + fig, axes = plt.subplots(1, len(early_values) + 1, figsize=(5.0 * (len(early_values) + 1), 4.6), dpi=170) + + panels = [("fixed K0", "fixed_gap", df)] + for early in early_values: + panels.append((f"linear velocity s={early}", "linear_gap", df[df["early_steps"] == early])) + + values = [] + for _title, column, group in panels: + values.extend(group[column].tolist()) + values.extend(group["empirical_gap"].tolist()) + lo, hi = float(min(values)), float(max(values)) + pad = 0.04 * (hi - lo + 1e-12) + + for ax, (title, column, group) in zip(axes, panels): + ax.scatter(group[column], group["empirical_gap"], s=26, alpha=0.72, color="#1f5a9d") + ax.plot([lo - pad, hi + pad], [lo - pad, hi + pad], color="black", linewidth=1.0) + err = group[column] - group["empirical_gap"] + ax.set_title(f"{title}\nMAE={err.abs().mean():.4f}, bias={err.mean():+.4f}") + ax.set_xlabel("predicted gap") + ax.set_ylabel("empirical gap") + ax.grid(alpha=0.18) + ax.set_xlim(lo - pad, hi + pad) + ax.set_ylim(lo - pad, hi + pad) + + fig.suptitle("First-order tangent-hierarchy predictor", y=1.02) + fig.tight_layout() + path = outdir / "first_order_prediction_scatter_by_early_step.png" + fig.savefig(path, bbox_inches="tight") + plt.close(fig) + return path + + +def plot_metrics(metric_df: pd.DataFrame, outdir: Path) -> Path: + fig, axes = plt.subplots(1, 3, figsize=(13.5, 4.4), dpi=170) + predictors = ["fixed K0", "linear velocity", "retangent", "compressed"] + colors = { + "fixed K0": "#777777", + "linear velocity": "#1f5a9d", + "retangent": "#2b8a3e", + "compressed": "#8a4fb0", + } + + for ax, metric_name in zip(axes, ["mae", "bias", "corr"]): + for predictor in predictors: + sub = metric_df[metric_df["predictor"] == predictor].sort_values("early_steps") + ax.plot( + sub["early_steps"], + sub[metric_name], + marker="o", + linewidth=1.8, + color=colors[predictor], + label=predictor, + ) + if metric_name == "bias": + ax.axhline(0.0, color="black", linewidth=1.0) + ax.set_xlabel("early operator step s") + ax.set_title(metric_name.upper()) + ax.grid(alpha=0.18) + axes[0].set_ylabel("metric value") + axes[0].legend(fontsize=8) + fig.suptitle("Prediction metrics vs early operator horizon", y=1.02) + fig.tight_layout() + path = outdir / "first_order_error_metrics_by_early_step.png" + fig.savefig(path, bbox_inches="tight") + plt.close(fig) + return path + + +def plot_distributions(df: pd.DataFrame, outdir: Path) -> Path: + early_values = sorted(df["early_steps"].unique()) + fig, axes = plt.subplots(1, len(early_values), figsize=(5.0 * len(early_values), 4.4), dpi=170, squeeze=False) + for ax, early in zip(axes.ravel(), early_values): + group = df[df["early_steps"] == early] + bins = np.linspace( + min(group["empirical_gap"].min(), group["linear_gap"].min(), group["fixed_gap"].min()), + max(group["empirical_gap"].max(), group["linear_gap"].max(), group["fixed_gap"].max()), + 18, + ) + ax.hist(group["empirical_gap"], bins=bins, alpha=0.45, density=True, color="#c65f16", label="empirical") + ax.hist(group["fixed_gap"], bins=bins, alpha=0.32, density=True, color="#777777", label="fixed K0") + ax.hist(group["linear_gap"], bins=bins, alpha=0.36, density=True, color="#1f5a9d", label="linear velocity") + ax.set_title(f"gap distribution, s={early}") + ax.set_xlabel("FA/BP train gap") + ax.set_ylabel("density") + ax.legend(fontsize=8) + ax.grid(alpha=0.18) + fig.tight_layout() + path = outdir / "first_order_gap_distribution_overlay.png" + fig.savefig(path, bbox_inches="tight") + plt.close(fig) + return path + + +def main() -> None: + args = parse_args() + args.outdir.mkdir(parents=True, exist_ok=True) + df = pd.read_csv(args.csv) + metric_df = metrics(df) + metric_path = args.outdir / "first_order_metrics.csv" + metric_df.to_csv(metric_path, index=False) + paths = [ + plot_scatter(df, args.outdir), + plot_metrics(metric_df, args.outdir), + plot_distributions(df, args.outdir), + ] + print(f"metrics: {metric_path}") + for path in paths: + print(f"plot: {path}") + print(metric_df.to_string(index=False)) + + +if __name__ == "__main__": + main() -- cgit v1.2.3