#!/usr/bin/env python3 """Small author-code EP endpoint for structured-measurement-bias screening.""" from __future__ import annotations import argparse import json import os from pathlib import Path import random import subprocess import sys import time import numpy as np import torch ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from sdil.rain_ep_adapter import ( # noqa: E402 RainGradientCorrector, RainLayerStateCorrector, attach_layer_to_rain_estimator, attach_to_rain_estimator, observe_rain_neutral, ) PINNED_REVISION = "6b253fd8a5d267535f58ab79992256ef10031ceb" def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--author-root", type=Path, required=True) parser.add_argument("--device", default="cuda") parser.add_argument( "--adapter", choices=("parameter", "layer"), default="parameter") parser.add_argument( "--beta-policy", choices=("fixed_positive", "random_sign"), default="fixed_positive") parser.add_argument("--beta-seed", type=int, default=7100) parser.add_argument( "--mode", choices=sorted(RainGradientCorrector.MODES), required=True) parser.add_argument("--bias-ratio", type=float, default=0.5) parser.add_argument("--predictor-rate", type=float, default=0.1) parser.add_argument( "--neutral-cadence", type=int, default=1, help="training steps per neutral update; zero freezes after calibration") parser.add_argument("--calibration-batches", type=int, default=0) parser.add_argument("--layer-calibration-steps", type=int, default=1) parser.add_argument( "--layer-bias-normalization", choices=("clean_difference", "first_state"), default="clean_difference") parser.add_argument("--epochs", type=int, default=2) parser.add_argument("--train-limit", type=int, default=2048) parser.add_argument("--test-limit", type=int, default=1024) parser.add_argument( "--evaluation-split", choices=("test", "train_holdout"), default="test") parser.add_argument("--data-seed", type=int, default=1988) parser.add_argument("--batch-size", type=int, default=128) parser.add_argument("--training-iterations", type=int, default=12) parser.add_argument("--inference-iterations", type=int, default=30) parser.add_argument("--deterministic", action="store_true") parser.add_argument("--seed", type=int, default=1988) parser.add_argument("--output", type=Path, required=True) return parser.parse_args() def revision(path: Path) -> str: return subprocess.check_output( ["git", "-C", str(path), "rev-parse", "HEAD"], text=True).strip() def accuracy(cost, size: int) -> float: return float((~cost.error_fn()).float().sum()) / size @torch.no_grad() def evaluate(network, cost, minimizer, loader) -> tuple[float, float]: total_correct = 0.0 total_cost = 0.0 total = 0 for x, y in loader: network.set_input(x, reset=True) minimizer.compute_equilibrium() cost.set_target(y) batch = x.shape[0] total_correct += accuracy(cost, batch) * batch total_cost += float(cost.eval().sum()) total += batch return total_correct / total, total_cost / total def main() -> None: args = parse_args() if args.calibration_batches < 0: raise ValueError("calibration batches must be nonnegative") if args.beta_policy == "random_sign" and not ( args.adapter == "layer" and args.mode == "raw" ): raise ValueError( "random-sign beta is a raw layer-measurement baseline") author_root = args.author_root.resolve() author_revision = revision(author_root) if author_revision != PINNED_REVISION: raise ValueError( f"expected Rain revision {PINNED_REVISION}, got {author_revision}") sys.path.insert(0, str(author_root)) from datasets import load_fashion_mnist from model.function.cost import SquaredError from model.function.network import Network from model.hopfield.minimizer import FixedPointMinimizer from model.hopfield.network import ConvHopfieldEnergy28 from training.sgd import AugmentedFunction, EquilibriumProp random.seed(args.seed) np.random.seed(args.seed) torch.manual_seed(args.seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(args.seed) if args.deterministic: torch.backends.cudnn.benchmark = False torch.backends.cudnn.deterministic = True torch.use_deterministic_algorithms(True, warn_only=True) device = torch.device(args.device) training_data, test_data = load_fashion_mnist( normalize=True, augment_32x32=False) split_generator = torch.Generator().manual_seed(args.data_seed) split_indices = torch.randperm( len(training_data), generator=split_generator) if args.evaluation_split == "train_holdout": if args.train_limit + args.test_limit > len(training_data): raise ValueError("training and holdout subsets overlap") train_indices = split_indices[:args.train_limit] evaluation_data = training_data evaluation_indices = split_indices[ args.train_limit:args.train_limit + args.test_limit] else: train_indices = split_indices[:args.train_limit] evaluation_data = test_data evaluation_indices = torch.arange(min(args.test_limit, len(test_data))) train_generator = torch.Generator().manual_seed(args.seed + 65537) training_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(training_data, train_indices.tolist()), batch_size=args.batch_size, shuffle=True, generator=train_generator, num_workers=0) calibration_generator = torch.Generator().manual_seed(args.seed + 104729) calibration_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(training_data, train_indices.tolist()), batch_size=args.batch_size, shuffle=True, generator=calibration_generator, num_workers=0) test_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(evaluation_data, evaluation_indices.tolist()), batch_size=args.batch_size, shuffle=False, num_workers=0) energy = ConvHopfieldEnergy28( num_inputs=1, num_hiddens_1=32, num_hiddens_2=64, num_outputs=10, weight_gains=[0.6, 0.6, 1.5]) energy.set_device(str(device)) network = Network(energy) cost = SquaredError(energy.layers()[-1]) augmented = AugmentedFunction(energy, cost) training_minimizer = FixedPointMinimizer( augmented, network.free_layers()) training_minimizer.mode = "asynchronous" training_minimizer.num_iterations = args.training_iterations estimator = EquilibriumProp( energy.params(), energy.layers(), augmented, cost, training_minimizer) estimator.variant = "positive" estimator.nudging = 0.25 if args.adapter == "parameter": corrector = RainGradientCorrector( mode=args.mode, bias_ratio=args.bias_ratio, predictor_rate=args.predictor_rate, neutral_cadence=args.neutral_cadence, seed=args.seed + 1729) attach_to_rain_estimator(estimator, corrector) else: if args.calibration_batches: raise ValueError( "layer adapter calibrates inside existing free phases; " "external calibration batches must be zero") corrector = RainLayerStateCorrector( mode=args.mode, bias_ratio=args.bias_ratio, predictor_rate=args.predictor_rate, calibration_steps=args.layer_calibration_steps, bias_normalization=args.layer_bias_normalization, seed=args.seed + 1729) attach_layer_to_rain_estimator(estimator, corrector) inference_minimizer = FixedPointMinimizer( energy, network.free_layers()) inference_minimizer.mode = "asynchronous" inference_minimizer.num_iterations = args.inference_iterations learning_rates = [0.02] * len(energy.params()) parameter_groups = [ {"params": parameter.state, "lr": learning_rate} for parameter, learning_rate in zip(energy.params(), learning_rates) ] optimizer = torch.optim.SGD( parameter_groups, lr=0.1, momentum=0.9, weight_decay=3e-4) metrics = [] start = time.time() calibration_start = time.time() calibration_observations = 0 if args.calibration_batches: if args.adapter != "parameter": raise AssertionError("layer calibration was not rejected above") if args.mode not in {"constant", "innovation"}: raise ValueError( "precalibration is defined only for constant or innovation mode") while calibration_observations < args.calibration_batches: for x, _ in calibration_loader: network.set_input(x, reset=False) inference_minimizer.compute_equilibrium() observe_rain_neutral(estimator, corrector) calibration_observations += 1 if calibration_observations == args.calibration_batches: break calibration_seconds = time.time() - calibration_start beta_generator = torch.Generator().manual_seed(args.beta_seed) beta_sign_counts = {"positive": 0, "negative": 0} for epoch in range(1, args.epochs + 1): total_cost = 0.0 total_correct = 0.0 total = 0 for x, y in training_loader: network.set_input(x, reset=False) inference_minimizer.compute_equilibrium() cost.set_target(y) batch = x.shape[0] total_cost += float(cost.eval().sum()) total_correct += accuracy(cost, batch) * batch total += batch if args.beta_policy == "random_sign": sign = 1 if int(torch.randint( 0, 2, (), generator=beta_generator)) else -1 estimator._first_nudging = 0.0 estimator._second_nudging = sign * estimator.nudging beta_sign_counts[ "positive" if sign > 0 else "negative"] += 1 gradients = estimator.compute_gradient() if any(gradient.requires_grad for gradient in gradients): raise AssertionError("adapter produced a requires-grad tensor") for parameter, gradient in zip(energy.params(), gradients): parameter.state.grad = gradient optimizer.step() for parameter in energy.params(): parameter.clamp_() test_accuracy, test_cost = evaluate( network, cost, inference_minimizer, test_loader) finite = all( bool(torch.isfinite(parameter.state).all()) for parameter in energy.params() ) record = { "epoch": epoch, "train_accuracy": total_correct / total, "train_cost": total_cost / total, "test_accuracy": test_accuracy, "test_cost": test_cost, "finite": finite, "corrector": dict(corrector.last_diagnostics), "wall_seconds": time.time() - start, } metrics.append(record) print(json.dumps(record), flush=True) if not finite: break report = { "schema": "rain_ep_structured_bias_screen_v1", "sdil": {"revision": revision(ROOT)}, "author": { "repository": "https://github.com/rain-neuromorphics/energy-based-learning", "revision": author_revision, }, "protocol": { "dataset": "FashionMNIST", "network": "author ConvHopfieldEnergy28 32-64-10", "algorithm": "equilibrium propagation", "beta_policy": args.beta_policy, "beta_seed": args.beta_seed, "adapter": args.adapter, "mode": args.mode, "bias_ratio": args.bias_ratio, "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, "layer_calibration_steps": args.layer_calibration_steps, "layer_bias_normalization": args.layer_bias_normalization, "calibration_batches": args.calibration_batches, "calibration_observations": calibration_observations, "calibration_seconds": calibration_seconds, "extra_equilibrium_phases_for_predictor": ( 0 if args.adapter == "layer" else args.calibration_batches), "predictor_neutral_source": ( "existing_first_EP_phase" if args.adapter == "layer" else "separate_free_equilibrium"), "bias_ratio_normalization": ( ( "initial_free_layer_state_rms" if args.layer_bias_normalization == "first_state" else "experimenter_initial_clean_layer_state_difference_rms" ) if args.adapter == "layer" else "initial_local_parameter_state_rms"), "bias_ratio_normalization_visible_to_predictor": False, "epochs": args.epochs, "train_limit": args.train_limit, "test_limit": args.test_limit, "evaluation_split": args.evaluation_split, "data_seed": args.data_seed, "batch_size": args.batch_size, "training_iterations": args.training_iterations, "inference_iterations": args.inference_iterations, "seed": args.seed, "device": str(device), "determinism": ( "best_effort_warn_only" if args.deterministic else "author_default"), "autodiff_used_for_learning": False, }, "hardware": { "torch_version": torch.__version__, "torch_cuda_version": torch.version.cuda, "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"), "device_name": ( torch.cuda.get_device_name(device) if device.type == "cuda" else "cpu"), }, "metrics": metrics, "epochs_completed": len(metrics), "beta_sign_counts": beta_sign_counts, "final": metrics[-1], } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(report, indent=2) + "\n") print(f"wrote {args.output}") if __name__ == "__main__": main()