summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:40:47 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:40:47 -0500
commit843d9c33d7dbdef8f2ebc88468f3736154a91004 (patch)
treefeb8bbe5407b3dd13a746e70756d4b78ff03b735
parent5781e8678f2406f63c1cadfe2772c29e4b7ed968 (diff)
fix: avoid tensorboard dependency in Rain endpoint
-rw-r--r--experiments/rain_ep_bias_train.py10
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()