From 2d266de3e641081988f279c45b8ef32364fa0cf5 Mon Sep 17 00:00:00 2001 From: Jake VanderPlas Date: Thu, 6 Aug 2026 08:00:02 -0700 Subject: [PATCH] No public description PiperOrigin-RevId: 960303794 --- learned_optimization/learned_optimizers/nn_adam.py | 8 ++++---- learned_optimization/tasks/es_wrapper.py | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/learned_optimization/learned_optimizers/nn_adam.py b/learned_optimization/learned_optimizers/nn_adam.py index dad022f..b3a01ec 100644 --- a/learned_optimization/learned_optimizers/nn_adam.py +++ b/learned_optimization/learned_optimizers/nn_adam.py @@ -231,13 +231,13 @@ def init(self, key: PRNGKey) -> lopt_base.MetaParams: self.rnn_to_controls.init(key3, jnp.zeros([0, self.lstm_hidden_size])), "per_layer_lr": - _scaled_lr.forward(self.initial_learning_rate), + _scaled_lr.forward(self.initial_learning_rate), # pyrefly: ignore[bad-argument-type] "per_layer_beta1": - _scaled_one_minus_log.forward(self.initial_beta1), + _scaled_one_minus_log.forward(self.initial_beta1), # pyrefly: ignore[bad-argument-type] "per_layer_beta2": - _scaled_one_minus_log.forward(self.initial_beta2), + _scaled_one_minus_log.forward(self.initial_beta2), # pyrefly: ignore[bad-argument-type] "per_layer_epsilon": - _scaled_epsilon.forward(self.initial_epsilon), + _scaled_epsilon.forward(self.initial_epsilon), # pyrefly: ignore[bad-argument-type] }) def opt_fn(self, diff --git a/learned_optimization/tasks/es_wrapper.py b/learned_optimization/tasks/es_wrapper.py index 870366d..31541e1 100644 --- a/learned_optimization/tasks/es_wrapper.py +++ b/learned_optimization/tasks/es_wrapper.py @@ -52,7 +52,7 @@ def _fn(key: PRNGKey) -> Tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray]: pos = _sample_perturbations(theta, key, std=std) p_theta = jax.tree_util.tree_map(lambda t, a: t + a, theta, pos) n_theta = jax.tree_util.tree_map(lambda t, a: t - a, theta, pos) - return pos, p_theta, n_theta + return pos, p_theta, n_theta # pyrefly: ignore[bad-return] keys = jax.random.split(key, num_samples) vec_pos, vec_p_theta, vec_n_theta = jax.vmap(_fn)(keys)