Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Cyclical SGLD

In this example we will demonstrate how Blackjax can be used to create non-trivial samplers by implementing Cyclical SGLD Zhang et al., 2019. Stochastic Gradient MCMC algorithms are typically used to sample from the posterior distribution of Bayesian Neural Networks. They differ from other gradient-based MCMC algorithms in that they estimate the gradient with minibatches of data instead of the full dataset.

However, SGMCMC algorithms are inefficient at exploring multimodal distributions which are typical of neural networks. To see this let’s consider a simple yet challenging example, an array of 25 gaussian distributions:

Notebook Cell
Source
<Figure size 800x800 with 1 Axes>

SGLD

Let us build a SGLD sampler with Blackjax with a decreading learning rate, and generate samples from this distribution:

Loading...
Loading...

As one can see on the following figure, SGLD has a hard time escaping the mode in which it started, leading to a poor approximation of the distribution:

Source
<Figure size 800x800 with 1 Axes>

Cyclical SGLD

To escape modes and better explore distributions, Cyclical SLGD alternes between two phases:

  1. Exploration using Stochastic Gradient Descent with a large-ish step size;

  2. Sampling using SGLD with a lower learning rate.

Cyclical schedule

Both the step size and the phase the algorithm is in are governed by a cyclical schedule which is built as follows:

Let us visualize a schedule for 200k training steps divided in 4 cycles. At each cycle 1/4th of the steps are dedicated to exploration:

Source
<Figure size 1200x400 with 1 Axes>

Step function

Let us now build a step function for the Cyclical SGLD algorithm that can act as a drop-in replacement to the SGLD kernel.

We leave the implementation of Cyclical SGHMC as an exercise for the reader.

Let’s sample using Cyclical SGLD, for the same number of iterations as with SGLD. We’ll use 30 cycles, and sample 75% of the time.

Loading...

By looking at the trajectory of the sampler it seems that the distribution is much better explored:

Source
<Figure size 800x800 with 1 Axes>

And the distribution indeed looks more better:

Source
<Figure size 800x800 with 1 Axes>
References
  1. Zhang, R., Li, C., Zhang, J., Chen, C., & Wilson, A. G. (2019). Cyclical Stochastic Gradient MCMC for Bayesian Deep Learning. International Conference on Learning Representations.