#!/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()