From 11375078d783c4fa47a1901aac297c7c52e9c9d9 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 3 Jun 2026 10:02:51 -0500 Subject: Add early-kernel operator predictor --- scripts/compressed_operator_predictor.py | 389 +++++++++++++++++++++++++++++++ 1 file changed, 389 insertions(+) create mode 100644 scripts/compressed_operator_predictor.py (limited to 'scripts/compressed_operator_predictor.py') diff --git a/scripts/compressed_operator_predictor.py b/scripts/compressed_operator_predictor.py new file mode 100644 index 0000000..ce52c7a --- /dev/null +++ b/scripts/compressed_operator_predictor.py @@ -0,0 +1,389 @@ +#!/usr/bin/env python3 +"""Compressed-operator finite-time predictor for FA/BP gaps.""" + +from __future__ import annotations + +import argparse +import csv +import json +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 fa_tangent_kernel_capacity import pseudo_jacobian # noqa: E402 + + +@dataclass(frozen=True) +class CompressedRow: + init_seed: int + feedback_seed: int + target_steps: int + early_steps: int + empirical_bp_loss: float + empirical_fa_loss: float + empirical_gap: float + fixed_bp_loss: float + fixed_fa_loss: float + fixed_gap: float + compressed_fa_loss: float + compressed_gap: float + linear_bp_loss: float + linear_fa_loss: float + linear_gap: float + retangent_bp_loss: float + retangent_fa_loss: float + retangent_gap: float + alpha_s: float + alpha_slope: float + fixed_gap_error: float + compressed_gap_error: float + linear_gap_error: float + retangent_gap_error: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Compressed operator predictor.") + 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("--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("--early-steps", type=int, nargs="+", default=[1, 2, 5]) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--init-seeds", type=int, default=2) + parser.add_argument("--feedback-seeds", type=int, default=4) + 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/compressed_operator_predictor"), + ) + 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 loss_from_residual(residual: np.ndarray, samples: int) -> float: + return 0.5 * float(residual @ residual) / samples + + +def kernel_pair( + weights: list[torch.Tensor], + x: torch.Tensor, + feedback: list[torch.Tensor], +) -> tuple[np.ndarray, np.ndarray]: + j_bp = pseudo_jacobian(weights, x, feedback=None).cpu().numpy() + j_fa = pseudo_jacobian(weights, x, feedback=feedback).cpu().numpy() + return j_bp @ j_bp.T, j_bp @ j_fa.T + + +def sgd_step( + weights: list[torch.Tensor], + x: torch.Tensor, + y: torch.Tensor, + lr: float, + feedback: list[torch.Tensor] | None, +) -> list[torch.Tensor]: + grads = dcs.gradients(weights, x, y, feedback) + return [weight - lr * grad for weight, grad in zip(weights, grads)] + + +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): + current = sgd_step(current, x, y, lr, feedback) + return current + + +def fixed_rollout( + kernel: np.ndarray, + residual: np.ndarray, + lr: float, + steps: int, + samples: int, +) -> np.ndarray: + current = residual.copy() + scale = lr / samples + for _ in range(steps): + current = current - scale * (kernel @ current) + return current + + +def compressed_rollout( + k_fa0: np.ndarray, + k_bp0: np.ndarray, + alpha_slope: float, + residual: np.ndarray, + lr: float, + steps: int, + samples: int, +) -> np.ndarray: + current = residual.copy() + scale = lr / samples + direction = k_bp0 - k_fa0 + for t in range(steps): + alpha_t = float(np.clip(alpha_slope * t, 0.0, 1.0)) + kernel = k_fa0 + alpha_t * direction + current = current - scale * (kernel @ current) + return current + + +def linear_velocity_rollout( + kernel0: np.ndarray, + kernel_s: np.ndarray, + early_step: int, + residual: np.ndarray, + lr: float, + steps: int, + samples: int, +) -> np.ndarray: + current = residual.copy() + scale = lr / samples + velocity = (kernel_s - kernel0) / early_step + for t in range(steps): + kernel = kernel0 + t * velocity + current = current - scale * (kernel @ current) + return current + + +def projection_alpha(k_fa_s: np.ndarray, k_fa0: np.ndarray, k_bp0: np.ndarray) -> float: + direction = k_bp0 - k_fa0 + denom = float(np.sum(direction * direction)) + if denom <= 0: + return 0.0 + alpha = float(np.sum((k_fa_s - k_fa0) * direction) / denom) + return float(np.clip(alpha, 0.0, 1.0)) + + +def run_one( + config: dcs.RunConfig, + x: torch.Tensor, + y: torch.Tensor, + init_seed: int, + feedback_seed: int, + early_step: int, +) -> CompressedRow: + initial = dcs.initialize_weights(config, config.widths[0], init_seed) + feedback = dcs.init_feedback(config, config.widths[0], feedback_seed) + with torch.no_grad(): + r0 = (dcs.predict(initial, x) - y).reshape(-1).cpu().numpy() + + k_bp0, k_fa0 = kernel_pair(initial, x, feedback) + fixed_bp_residual = fixed_rollout(k_bp0, r0, config.lr, config.steps, config.train_samples) + fixed_fa_residual = fixed_rollout(k_fa0, r0, config.lr, config.steps, config.train_samples) + + fa_early = train_to_steps(initial, x, y, config.lr, early_step, feedback=feedback) + _kbp_unused, k_fa_s = kernel_pair(fa_early, x, feedback) + alpha_s = projection_alpha(k_fa_s, k_fa0, k_bp0) + alpha_slope = alpha_s / early_step + compressed_fa_residual = compressed_rollout( + k_fa0, + k_bp0, + alpha_slope, + r0, + config.lr, + config.steps, + config.train_samples, + ) + + bp_early = train_to_steps(initial, x, y, config.lr, early_step, feedback=None) + k_bp_s, _kfa_unused = kernel_pair(bp_early, x, feedback) + linear_bp_residual = linear_velocity_rollout( + k_bp0, + k_bp_s, + early_step, + r0, + config.lr, + config.steps, + config.train_samples, + ) + linear_fa_residual = linear_velocity_rollout( + k_fa0, + k_fa_s, + early_step, + r0, + config.lr, + config.steps, + config.train_samples, + ) + with torch.no_grad(): + bp_early_residual = (dcs.predict(bp_early, x) - y).reshape(-1).cpu().numpy() + fa_early_residual = (dcs.predict(fa_early, x) - y).reshape(-1).cpu().numpy() + retangent_bp_residual = fixed_rollout( + k_bp_s, + bp_early_residual, + config.lr, + config.steps - early_step, + config.train_samples, + ) + retangent_fa_residual = fixed_rollout( + k_fa_s, + fa_early_residual, + config.lr, + config.steps - early_step, + 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_bp = dcs.mse(bp_target, x, y) + empirical_fa = dcs.mse(fa_target, x, y) + fixed_bp = loss_from_residual(fixed_bp_residual, config.train_samples) + fixed_fa = loss_from_residual(fixed_fa_residual, config.train_samples) + compressed_fa = loss_from_residual(compressed_fa_residual, config.train_samples) + linear_bp = loss_from_residual(linear_bp_residual, config.train_samples) + linear_fa = loss_from_residual(linear_fa_residual, config.train_samples) + retangent_bp = loss_from_residual(retangent_bp_residual, config.train_samples) + retangent_fa = loss_from_residual(retangent_fa_residual, config.train_samples) + empirical_gap = empirical_fa - empirical_bp + fixed_gap = fixed_fa - fixed_bp + compressed_gap = compressed_fa - fixed_bp + linear_gap = linear_fa - linear_bp + retangent_gap = retangent_fa - retangent_bp + + return CompressedRow( + init_seed=init_seed, + feedback_seed=feedback_seed, + target_steps=config.steps, + early_steps=early_step, + empirical_bp_loss=empirical_bp, + empirical_fa_loss=empirical_fa, + empirical_gap=empirical_gap, + fixed_bp_loss=fixed_bp, + fixed_fa_loss=fixed_fa, + fixed_gap=fixed_gap, + compressed_fa_loss=compressed_fa, + compressed_gap=compressed_gap, + linear_bp_loss=linear_bp, + linear_fa_loss=linear_fa, + linear_gap=linear_gap, + retangent_bp_loss=retangent_bp, + retangent_fa_loss=retangent_fa, + retangent_gap=retangent_gap, + alpha_s=alpha_s, + alpha_slope=alpha_slope, + fixed_gap_error=fixed_gap - empirical_gap, + compressed_gap_error=compressed_gap - empirical_gap, + linear_gap_error=linear_gap - empirical_gap, + retangent_gap_error=retangent_gap - empirical_gap, + ) + + +def write_rows(path: Path, rows: list[CompressedRow]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(CompressedRow.__annotations__.keys())) + writer.writeheader() + for row in rows: + writer.writerow(asdict(row)) + + +def main() -> None: + args = parse_args() + if args.torch_threads > 0: + torch.set_num_threads(args.torch_threads) + config = make_config(args) + x_train, y_train, _x_test, _y_test, _x_probe, _teacher = dcs.make_data(config) + + rows: list[CompressedRow] = [] + total = args.init_seeds * args.feedback_seeds * len(args.early_steps) + 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 + for early_step in args.early_steps: + count += 1 + print( + f"[{count}/{total}] init={init_seed} feedback={feedback_seed} " + f"early={early_step}", + flush=True, + ) + rows.append( + run_one(config, x_train, y_train, init_seed, feedback_seed, early_step) + ) + + outdir = Path(args.outdir) + outdir.mkdir(parents=True, exist_ok=True) + write_rows(outdir / "compressed_operator_rows.csv", rows) + payload = { + "config": { + key: str(value) if isinstance(value, Path) else value + for key, value in vars(args).items() + }, + "rows": [asdict(row) for row in rows], + } + (outdir / "summary.json").write_text(json.dumps(payload, indent=2) + "\n") + fixed = np.array([row.fixed_gap_error for row in rows]) + compressed = np.array([row.compressed_gap_error for row in rows]) + linear = np.array([row.linear_gap_error for row in rows]) + retangent = np.array([row.retangent_gap_error for row in rows]) + print(f"rows: {outdir / 'compressed_operator_rows.csv'}") + print(f"fixed MAE={np.mean(np.abs(fixed)):.6g}, bias={fixed.mean():.6g}") + print( + f"compressed MAE={np.mean(np.abs(compressed)):.6g}, " + f"bias={compressed.mean():.6g}" + ) + print(f"linear MAE={np.mean(np.abs(linear)):.6g}, bias={linear.mean():.6g}") + print( + f"retangent MAE={np.mean(np.abs(retangent)):.6g}, " + f"bias={retangent.mean():.6g}" + ) + + +if __name__ == "__main__": + main() -- cgit v1.2.3