summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/rain_ep_bias_train.py23
1 files changed, 22 insertions, 1 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py
index 80135ea..6528be6 100644
--- a/experiments/rain_ep_bias_train.py
+++ b/experiments/rain_ep_bias_train.py
@@ -37,6 +37,10 @@ def parse_args() -> argparse.Namespace:
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)
@@ -90,6 +94,11 @@ 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:
@@ -212,6 +221,8 @@ def main() -> None:
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
@@ -224,6 +235,13 @@ def main() -> None:
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")
@@ -263,7 +281,9 @@ def main() -> None:
"protocol": {
"dataset": "FashionMNIST",
"network": "author ConvHopfieldEnergy28 32-64-10",
- "algorithm": "positive equilibrium propagation",
+ "algorithm": "equilibrium propagation",
+ "beta_policy": args.beta_policy,
+ "beta_seed": args.beta_seed,
"adapter": args.adapter,
"mode": args.mode,
"bias_ratio": args.bias_ratio,
@@ -306,6 +326,7 @@ def main() -> None:
},
"metrics": metrics,
"epochs_completed": len(metrics),
+ "beta_sign_counts": beta_sign_counts,
"final": metrics[-1],
}
args.output.parent.mkdir(parents=True, exist_ok=True)