diff --git a/optax/contrib/_schedule_free.py b/optax/contrib/_schedule_free.py index 8629e6c5d..a2a407400 100644 --- a/optax/contrib/_schedule_free.py +++ b/optax/contrib/_schedule_free.py @@ -147,7 +147,7 @@ def init_fn(params: base.Params) -> ScheduleFreeState: return ScheduleFreeState( b1=jnp.asarray(b1, dtype=params_dtype), weight_sum=jnp.zeros([], dtype=params_dtype), - step_count=jnp.ones([], dtype=jnp.int32), + step_count=jnp.zeros([], dtype=jnp.int32), max_lr=jnp.zeros([], dtype=params_dtype), base_optimizer_state=base_optimizer.init(params), z=z, diff --git a/optax/contrib/_schedule_free_test.py b/optax/contrib/_schedule_free_test.py index e78ac7d58..dc0254fac 100644 --- a/optax/contrib/_schedule_free_test.py +++ b/optax/contrib/_schedule_free_test.py @@ -40,6 +40,26 @@ def get_updates(params): class ScheduleFreeTest(parameterized.TestCase): + def test_step_count_starts_at_zero(self): + seen_counts = [] + + def learning_rate(count): + seen_counts.append(int(count)) + return 1.0 + + base_opt = alias.sgd(learning_rate=0.0, momentum=0.0) + opt = _schedule_free.schedule_free(base_opt, learning_rate=learning_rate) + params = jnp.array([1.0, 2.0], dtype=jnp.float32) + state = opt.init(params) + + self.assertEqual(int(state.step_count), 0) + + updates = jnp.zeros_like(params) + _, state = opt.update(updates, state, params) + + self.assertEqual(seen_counts, [0]) + self.assertEqual(int(state.step_count), 1) + def test_learning_rate_zero(self): base_opt = alias.sgd(learning_rate=0.0, momentum=0.0) opt = _schedule_free.schedule_free(base_opt, learning_rate=0.0)