diff options
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 35 |
1 files changed, 28 insertions, 7 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index 1ad476d..db6f668 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -37,6 +37,9 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--adapter", choices=("parameter", "layer"), default="parameter") parser.add_argument( + "--network-protocol", choices=("conv28_screen", "comparative32"), + default="conv28_screen") + parser.add_argument( "--beta-policy", choices=("fixed_positive", "fixed_negative", "random_sign"), default="fixed_positive") @@ -64,6 +67,7 @@ def parse_args() -> argparse.Namespace: 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("--schedule-epochs", type=int, default=0) parser.add_argument("--deterministic", action="store_true") parser.add_argument("--seed", type=int, default=1988) parser.add_argument("--output", type=Path, required=True) @@ -114,7 +118,7 @@ def main() -> None: 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 model.hopfield.network import ConvHopfieldEnergy28, ConvHopfieldEnergy32 from training.sgd import AugmentedFunction, EquilibriumProp random.seed(args.seed) @@ -129,7 +133,7 @@ def main() -> None: device = torch.device(args.device) training_data, test_data = load_fashion_mnist( - normalize=True, augment_32x32=False) + normalize=True, augment_32x32=args.network_protocol == "comparative32") split_generator = torch.Generator().manual_seed(args.data_seed) split_indices = torch.randperm( len(training_data), generator=split_generator) @@ -158,9 +162,18 @@ def main() -> None: torch.utils.data.Subset(evaluation_data, evaluation_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]) + if args.network_protocol == "conv28_screen": + energy = ConvHopfieldEnergy28( + num_inputs=1, num_hiddens_1=32, num_hiddens_2=64, + num_outputs=10, weight_gains=[0.6, 0.6, 1.5]) + network_description = "author ConvHopfieldEnergy28 32-64-10" + learning_rates = [0.02] * len(energy.params()) + else: + energy = ConvHopfieldEnergy32( + num_inputs=1, num_outputs=10, weight_gains=[0.5] * 5) + network_description = "author comparative-study ConvHopfieldEnergy32" + layer_rates = [0.0625, 0.0375, 0.025, 0.02, 0.0125] + learning_rates = layer_rates + layer_rates energy.set_device(str(device)) network = Network(energy) cost = SquaredError(energy.layers()[-1]) @@ -200,13 +213,17 @@ def main() -> None: 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) + scheduler = ( + torch.optim.lr_scheduler.CosineAnnealingLR( + optimizer, T_max=args.schedule_epochs, eta_min=2e-6) + if args.schedule_epochs else None + ) metrics = [] start = time.time() @@ -280,6 +297,8 @@ def main() -> None: } metrics.append(record) print(json.dumps(record), flush=True) + if scheduler is not None: + scheduler.step() if not finite: break @@ -292,7 +311,8 @@ def main() -> None: }, "protocol": { "dataset": "FashionMNIST", - "network": "author ConvHopfieldEnergy28 32-64-10", + "network": network_description, + "network_protocol": args.network_protocol, "algorithm": "equilibrium propagation", "beta_policy": args.beta_policy, "beta_seed": args.beta_seed, @@ -327,6 +347,7 @@ def main() -> None: "batch_size": args.batch_size, "training_iterations": args.training_iterations, "inference_iterations": args.inference_iterations, + "schedule_epochs": args.schedule_epochs, "seed": args.seed, "device": str(device), "determinism": ( |
