diff --git a/docs/api/contrib.rst b/docs/api/contrib.rst index c30e3143e..db63afe2d 100644 --- a/docs/api/contrib.rst +++ b/docs/api/contrib.rst @@ -47,6 +47,9 @@ are not supported by the main library. ScheduleFreeState sophia SophiaState + spsa + spsa_gradient + SPSAState split_real_and_imaginary SplitRealAndImaginaryState scale_by_ademamix diff --git a/optax/contrib/__init__.py b/optax/contrib/__init__.py index 35c032ef0..792520ee3 100644 --- a/optax/contrib/__init__.py +++ b/optax/contrib/__init__.py @@ -74,3 +74,6 @@ from optax.contrib._sophia import HutchinsonState from optax.contrib._sophia import sophia from optax.contrib._sophia import SophiaState +from optax.contrib._spsa import spsa +from optax.contrib._spsa import spsa_gradient +from optax.contrib._spsa import SPSAState diff --git a/optax/contrib/_spsa.py b/optax/contrib/_spsa.py new file mode 100644 index 000000000..cd4d97efa --- /dev/null +++ b/optax/contrib/_spsa.py @@ -0,0 +1,175 @@ +# Copyright 2024 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""SPSA: Simultaneous Perturbation Stochastic Approximation. + +A contributed implementation of the gradient-free SPSA method from +"Multivariate Stochastic Approximation Using a Simultaneous Perturbation +Gradient Approximation" (https://doi.org/10.1109/9.119632) by James C. Spall. + +The objective is supplied at update time through the ``obj_fn`` keyword +argument, following the same convention as +:func:`optax.contrib.hutchinson_estimator_diag_hessian`. +""" + +from typing import NamedTuple, Optional + +import jax +import jax.numpy as jnp +from optax._src import base +from optax._src import combine +from optax._src import numerics +from optax._src import transform +import optax.tree + + +class SPSAState(NamedTuple): + """State for the SPSA gradient estimator.""" + + count: jax.Array # shape=(), dtype=jnp.int32 + key: jax.Array + + +def spsa_gradient( + c: jax.typing.ArrayLike = 0.1, + gamma: jax.typing.ArrayLike = 0.101, + seed: Optional[jax.Array] = None, +) -> base.GradientTransformationExtraArgs: + r"""Estimates the gradient via Simultaneous Perturbation Stochastic Approx. + + SPSA is a gradient-free method: rather than differentiating the objective, it + estimates the whole gradient from only **two** objective evaluations per step, + regardless of the number of parameters. This is useful when the objective is + non-differentiable, only available as a black box, or expensive to + differentiate. + + At step :math:`k` a random perturbation vector :math:`\Delta_k` with i.i.d. + :math:`\pm 1` entries is drawn and the gradient estimate is + + .. math:: + \hat{g}_{k,i} = \frac{f(\theta + c_k \Delta_k) - f(\theta - c_k \Delta_k)} + {2\, c_k\, \Delta_{k,i}}, + \qquad c_k = \frac{c}{(k + 1)^\gamma}. + + Because :math:`\Delta_{k,i} \in \{-1, +1\}` we have + :math:`1 / \Delta_{k,i} = \Delta_{k,i}`, so the estimate is a scalar times + :math:`\Delta_k`. Its expectation matches the true gradient up to + :math:`O(c_k^2)`. The returned updates are this gradient estimate; compose it + with a step size (e.g. :func:`optax.scale_by_learning_rate`) or use the + :func:`optax.contrib.spsa` wrapper. + + The objective is supplied at update time via the ``obj_fn`` keyword argument, + exactly like :func:`optax.contrib.hutchinson_estimator_diag_hessian`. + ``obj_fn`` must take ``params`` as its only argument and return a scalar:: + + obj_fn = lambda params: loss_fn(params, batch) + grad_estimate, state = estimator.update(grads, state, params, obj_fn=obj_fn) + + Args: + c: Base perturbation size :math:`c`; must be positive. + gamma: Decay exponent for the perturbation size. Spall recommends ``0.101``. + seed: Optional PRNG key used to draw the perturbation vectors. + + Returns: + A :class:`optax.GradientTransformationExtraArgs`. + + References: + Spall, `Multivariate Stochastic Approximation Using a Simultaneous + Perturbation Gradient Approximation + `_, IEEE TAC, 1992. + + .. seealso:: :func:`optax.contrib.spsa` + """ + + def init_fn(params): + del params + key = seed if seed is not None else jax.random.PRNGKey(0) + return SPSAState(count=jnp.zeros([], jnp.int32), key=key) + + def update_fn(updates, state, params=None, obj_fn=None, **extra_args): + # Complies with the GradientTransformationExtraArgs signature but ignores + # the incoming ``updates`` (SPSA estimates its own gradient) and any other + # extra args. + del updates, extra_args + if params is None: + raise ValueError('params must be provided to the spsa update function.') + if obj_fn is None: + raise ValueError('obj_fn must be provided to the spsa update function.') + + key, subkey = jax.random.split(state.key) + perturbation = optax.tree.random_like( + subkey, params, jax.random.rademacher, dtype=jnp.float32 + ) + perturbation = optax.tree.cast( + perturbation, optax.tree.dtype(params, 'lowest') + ) + + step = jnp.asarray(state.count, jnp.float32) + 1.0 + ck = c / step**gamma + + params_plus = jax.tree.map(lambda p, d: p + ck * d, params, perturbation) + params_minus = jax.tree.map(lambda p, d: p - ck * d, params, perturbation) + delta_obj = (obj_fn(params_plus) - obj_fn(params_minus)) / (2.0 * ck) + + # 1 / Delta_i == Delta_i for Delta_i in {-1, +1}. + grad_estimate = jax.tree.map(lambda d: delta_obj * d, perturbation) + return grad_estimate, SPSAState( + count=numerics.safe_increment(state.count), key=key + ) + + return base.GradientTransformationExtraArgs(init_fn, update_fn) + + +def spsa( + learning_rate: base.ScalarOrSchedule, + c: jax.typing.ArrayLike = 0.1, + gamma: jax.typing.ArrayLike = 0.101, + seed: Optional[jax.Array] = None, +) -> base.GradientTransformationExtraArgs: + r"""The SPSA (gradient-free) optimizer. + + Combines the SPSA gradient estimate (:func:`optax.contrib.spsa_gradient`) with + a learning rate. Only objective *values* are used, so the objective need not + be differentiable. The objective is passed at update time through ``obj_fn`` + (a function of ``params`` returning a scalar); the incoming ``grads`` are + ignored and may be ``None``-like:: + + obj_fn = lambda params: loss_fn(params, batch) + opt = optax.contrib.spsa(learning_rate=0.1) + state = opt.init(params) + updates, state = opt.update(grads, state, params, obj_fn=obj_fn) + params = optax.apply_updates(params, updates) + + Args: + learning_rate: The step size :math:`a_k`, either fixed or a schedule. + Spall's classic decaying choice ``a / (A + k + 1) ** alpha`` can be built + with :func:`optax.polynomial_schedule`. + c: Base perturbation size; must be positive. + gamma: Decay exponent for the perturbation size. + seed: Optional PRNG key used to draw the perturbation vectors. + + Returns: + A :class:`optax.GradientTransformationExtraArgs`. + + References: + Spall, `Multivariate Stochastic Approximation Using a Simultaneous + Perturbation Gradient Approximation + `_, IEEE TAC, 1992. + + .. seealso:: :func:`optax.contrib.spsa_gradient` + """ + return combine.chain( + spsa_gradient(c=c, gamma=gamma, seed=seed), + transform.scale_by_learning_rate(learning_rate), + ) diff --git a/optax/contrib/_spsa_test.py b/optax/contrib/_spsa_test.py new file mode 100644 index 000000000..10928bada --- /dev/null +++ b/optax/contrib/_spsa_test.py @@ -0,0 +1,121 @@ +# Copyright 2024 DeepMind Technologies Limited. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for the SPSA optimizer in `optax.contrib._spsa`.""" + +from absl.testing import absltest +import jax +import jax.numpy as jnp +import numpy as np +from optax import apply_updates +from optax.contrib import _spsa + + +class SPSATest(absltest.TestCase): + + def test_spsa_minimizes_quadratic(self): + # SPSA only sees objective values, never gradients. + target = jnp.array([1.0, -2.0, 3.0, 0.5]) + obj_fn = lambda params: 0.5 * jnp.sum((params - target) ** 2) + + params = jnp.zeros_like(target) + opt = _spsa.spsa(learning_rate=0.2, c=0.1, seed=jax.random.key(0)) + state = opt.init(params) + + @jax.jit + def step(params, state): + # `grads` are ignored by SPSA; pass zeros to prove they are unused. + updates, state = opt.update( + jnp.zeros_like(params), state, params, obj_fn=obj_fn + ) + return apply_updates(params, updates), state + + initial_loss = obj_fn(params) + for _ in range(3000): + params, state = step(params, state) + + self.assertLess(float(obj_fn(params)), 1e-3) + self.assertLess(float(obj_fn(params)), float(initial_loss)) + np.testing.assert_allclose(params, target, atol=2e-2) + + def test_spsa_minimizes_quadratic_with_pytree_params(self): + target = {'w': jnp.array([2.0, -1.0]), 'b': jnp.array(0.5)} + obj_fn = lambda p: 0.5 * ( + jnp.sum((p['w'] - target['w']) ** 2) + (p['b'] - target['b']) ** 2 + ) + + params = {'w': jnp.zeros(2), 'b': jnp.zeros(())} + opt = _spsa.spsa(learning_rate=0.2, c=0.1, seed=jax.random.key(1)) + state = opt.init(params) + + @jax.jit + def step(params, state): + updates, state = opt.update(params, state, params, obj_fn=obj_fn) + return apply_updates(params, updates), state + + for _ in range(3000): + params, state = step(params, state) + + self.assertLess(float(obj_fn(params)), 1e-3) + + def test_gradient_estimate_is_unbiased(self): + # For f(x) = 0.5||x - t||^2 the SPSA estimate is exactly unbiased (no + # O(c^2) bias term), so averaging many estimates recovers the true + # gradient `x - t`. + target = jnp.array([0.7, -1.3, 2.1]) + obj_fn = lambda params: 0.5 * jnp.sum((params - target) ** 2) + point = jnp.array([1.0, 1.0, 1.0]) + true_grad = point - target + + estimator = _spsa.spsa_gradient(c=0.05, seed=jax.random.key(2)) + state = estimator.init(point) + + @jax.jit + def one_estimate(state): + grad, state = estimator.update(point, state, point, obj_fn=obj_fn) + return grad, state + + acc = jnp.zeros_like(point) + n = 20000 + for _ in range(n): + grad, state = one_estimate(state) + acc = acc + grad + mean_estimate = acc / n + + np.testing.assert_allclose(mean_estimate, true_grad, atol=5e-2) + + def test_update_requires_obj_fn_and_params(self): + estimator = _spsa.spsa_gradient() + params = jnp.ones(3) + state = estimator.init(params) + + with self.assertRaises(ValueError): + estimator.update(params, state, params) # missing obj_fn + with self.assertRaises(ValueError): + estimator.update(params, state, obj_fn=jnp.sum) # no params + + def test_state_count_increments(self): + estimator = _spsa.spsa_gradient() + params = jnp.ones(2) + state = estimator.init(params) + self.assertEqual(int(state.count), 0) + _, state = estimator.update( + params, state, params, obj_fn=jnp.sum + ) + self.assertEqual(int(state.count), 1) + self.assertEqual(state.count.dtype, jnp.int32) + + +if __name__ == '__main__': + absltest.main()