#!/usr/bin/env python3 """Early kernel-drift predictors for finite-time 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 PredictorRow: 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 directional_bp_loss: float directional_fa_loss: float directional_gap: float kfa_fro_drift: float kbp_fro_drift: float fa_bp_overlap_0: float fa_bp_overlap_s: float fa_bp_overlap_slope: float fa_directional_gain_s: float bp_directional_gain_s: float fa_directional_drift_s: float bp_directional_drift_s: float def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Compute early kernel predictors.") 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=4) parser.add_argument("--feedback-seeds", type=int, default=8) 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/early_kernel_predictors"), ) 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 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 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 fixed_prediction( k_bp0: np.ndarray, k_fa0: np.ndarray, residual: np.ndarray, lr: float, steps: int, samples: int, ) -> tuple[float, float]: scale = lr / samples rb = residual.copy() rf = residual.copy() for _ in range(steps): rb = rb - scale * (k_bp0 @ rb) rf = rf - scale * (k_fa0 @ rf) return 0.5 * float(rb @ rb) / samples, 0.5 * float(rf @ rf) / samples def loss_from_residual(residual: np.ndarray, samples: int) -> float: return 0.5 * float(residual @ residual) / samples def overlap(a: np.ndarray, b: np.ndarray) -> float: denom = np.linalg.norm(a, "fro") * np.linalg.norm(b, "fro") if denom == 0: return 0.0 return float(np.sum(a * b) / denom) def directional_gain(kernel: np.ndarray, residual: np.ndarray) -> float: denom = float(residual @ residual) if denom == 0: return 0.0 return float(residual @ (kernel @ residual) / denom) def directional_loss_prediction( residual: np.ndarray, lambda_0: float, lambda_s: float, early_step: int, lr: float, target_steps: int, samples: int, ) -> float: slope = (lambda_s - lambda_0) / early_step cumulative = target_steps * lambda_0 + 0.5 * target_steps * (target_steps - 1) * slope # Keep the early linear extrapolation from producing an unstable negative # integrated rate; this is a stability projection, not a fitted coefficient. cumulative = max(cumulative, 0.0) initial_loss = loss_from_residual(residual, samples) return initial_loss * float(np.exp(-(2.0 * lr / samples) * cumulative)) 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 run_one( config: dcs.RunConfig, x: torch.Tensor, y: torch.Tensor, init_seed: int, feedback_seed: int, early_step: int, ) -> PredictorRow: 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() kbp0, kfa0 = kernel_pair(initial, x, feedback) fixed_bp, fixed_fa = fixed_prediction( kbp0, kfa0, r0, config.lr, config.steps, 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) bp_early = train_to_steps(initial, x, y, config.lr, early_step, feedback=None) fa_early = train_to_steps(initial, x, y, config.lr, early_step, feedback=feedback) kbp_s, _kfa_unused = kernel_pair(bp_early, x, feedback) _kbp_unused, kfa_s = kernel_pair(fa_early, x, feedback) with torch.no_grad(): rb_s = (dcs.predict(bp_early, x) - y).reshape(-1).cpu().numpy() rf_s = (dcs.predict(fa_early, x) - y).reshape(-1).cpu().numpy() kfa_fro_drift = float(np.linalg.norm(kfa_s - kfa0, "fro") / np.linalg.norm(kfa0, "fro")) kbp_fro_drift = float(np.linalg.norm(kbp_s - kbp0, "fro") / np.linalg.norm(kbp0, "fro")) fa_bp_overlap_0 = overlap(kfa0, kbp0) fa_bp_overlap_s = overlap(kfa_s, kbp_s) fa_directional_gain_0 = directional_gain(kfa0, r0) fa_directional_gain_s = directional_gain(kfa_s, rf_s) bp_directional_gain_0 = directional_gain(kbp0, r0) bp_directional_gain_s = directional_gain(kbp_s, rb_s) directional_bp = directional_loss_prediction( r0, bp_directional_gain_0, bp_directional_gain_s, early_step, config.lr, config.steps, config.train_samples, ) directional_fa = directional_loss_prediction( r0, fa_directional_gain_0, fa_directional_gain_s, early_step, config.lr, config.steps, config.train_samples, ) return PredictorRow( 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_fa - empirical_bp, fixed_bp_loss=fixed_bp, fixed_fa_loss=fixed_fa, fixed_gap=fixed_fa - fixed_bp, directional_bp_loss=directional_bp, directional_fa_loss=directional_fa, directional_gap=directional_fa - directional_bp, kfa_fro_drift=kfa_fro_drift, kbp_fro_drift=kbp_fro_drift, fa_bp_overlap_0=fa_bp_overlap_0, fa_bp_overlap_s=fa_bp_overlap_s, fa_bp_overlap_slope=(fa_bp_overlap_s - fa_bp_overlap_0) / early_step, fa_directional_gain_s=fa_directional_gain_s, bp_directional_gain_s=bp_directional_gain_s, fa_directional_drift_s=(fa_directional_gain_s - fa_directional_gain_0) / max( abs(fa_directional_gain_0), 1e-12 ), bp_directional_drift_s=(bp_directional_gain_s - bp_directional_gain_0) / max( abs(bp_directional_gain_0), 1e-12 ), ) def write_rows(path: Path, rows: list[PredictorRow]) -> None: path.parent.mkdir(parents=True, exist_ok=True) with path.open("w", newline="") as handle: writer = csv.DictWriter(handle, fieldnames=list(PredictorRow.__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[PredictorRow] = [] 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 / "early_kernel_predictors.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") print(f"rows: {outdir / 'early_kernel_predictors.csv'}") if __name__ == "__main__": main()