summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 17:24:54 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 17:24:54 -0500
commitc868f021850d4f8a91a4eca97ad475aa6fafdc7d (patch)
tree015f1ec0e988d3b78c47775bb62f5ddacf007dc6 /experiments/rain_ep_bias_train.py
parent41555c80d365eb23801e21e7795e894eacadde4e (diff)
feat: add Rain comparative-study protocol
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
-rw-r--r--experiments/rain_ep_bias_train.py35
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": (