diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 16:40:47 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 16:40:47 -0500 |
| commit | 843d9c33d7dbdef8f2ebc88468f3736154a91004 (patch) | |
| tree | feb8bbe5407b3dd13a746e70756d4b78ff03b735 /experiments/rain_ep_bias_train.py | |
| parent | 5781e8678f2406f63c1cadfe2772c29e4b7ed968 (diff) | |
fix: avoid tensorboard dependency in Rain endpoint
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 10 |
1 files changed, 6 insertions, 4 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index eed5715..871d094 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -84,7 +84,6 @@ def main() -> None: 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) @@ -137,9 +136,12 @@ def main() -> None: 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) + 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) metrics = [] start = time.time() |
