diff options
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 23 |
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) |
