diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-06-05 16:03:17 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-06-05 16:03:17 -0500 |
| commit | 39ea1acbb5ad62043cb60c0f96b3bcb109607831 (patch) | |
| tree | a3a411b34f51f926142bb0251e4f6eea083e413e /scripts/fa_tangent_hierarchy_derivative_probe.py | |
| parent | a3f6c103678e0dcae3a682945c784a3f88d5039f (diff) | |
Test first order FA tangent hierarchy
Diffstat (limited to 'scripts/fa_tangent_hierarchy_derivative_probe.py')
| -rw-r--r-- | scripts/fa_tangent_hierarchy_derivative_probe.py | 278 |
1 files changed, 278 insertions, 0 deletions
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() |
