diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 213 |
1 files changed, 213 insertions, 0 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py new file mode 100644 index 0000000..eed5715 --- /dev/null +++ b/experiments/rain_ep_bias_train.py @@ -0,0 +1,213 @@ +#!/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, +) + + +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("--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() + 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.monitor import Optimizer + 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) + 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()) + optimizer = Optimizer( + energy, cost, learning_rates, momentum=0.9, + weight_decay=3e-4) + + metrics = [] + start = time.time() + 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, + "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() |
