#!/usr/bin/env python3 """Small author-code EP endpoint for structured-measurement-bias screening.""" from __future__ import annotations import argparse import json 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, 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( "--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) parser.add_argument("--calibration-batches", type=int, default=0) 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("--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("--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") 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) device = torch.device(args.device) training_data, test_data = load_fashion_mnist( normalize=True, augment_32x32=False) train_generator = torch.Generator().manual_seed(args.seed) train_indices = torch.randperm( len(training_data), generator=train_generator)[:args.train_limit] test_indices = torch.arange(min(args.test_limit, len(test_data))) 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(test_data, test_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 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) 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.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 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 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) record = { "epoch": epoch, "train_accuracy": total_correct / total, "train_cost": total_cost / total, "test_accuracy": test_accuracy, "test_cost": test_cost, "corrector": dict(corrector.last_diagnostics), "wall_seconds": time.time() - start, } metrics.append(record) print(json.dumps(record), flush=True) report = { "schema": "rain_ep_structured_bias_screen_v1", "author": { "repository": "https://github.com/rain-neuromorphics/energy-based-learning", "revision": author_revision, }, "protocol": { "dataset": "FashionMNIST", "network": "author ConvHopfieldEnergy28 32-64-10", "algorithm": "positive equilibrium propagation", "mode": args.mode, "bias_ratio": args.bias_ratio, "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, "calibration_batches": args.calibration_batches, "calibration_observations": calibration_observations, "calibration_seconds": calibration_seconds, "epochs": args.epochs, "train_limit": args.train_limit, "test_limit": args.test_limit, "batch_size": args.batch_size, "training_iterations": args.training_iterations, "inference_iterations": args.inference_iterations, "seed": args.seed, "device": str(device), "autodiff_used_for_learning": False, }, "metrics": metrics, "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()