summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
-rw-r--r--experiments/rain_ep_bias_train.py213
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()