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