Variational Neural Inference: an executable tutorial

22 September 2026

Most introductions to variational inference either stop at the ELBO or jump straight to a published model with ten thousand lines of scaffolding. This tutorial is my attempt at the path between the two: eleven notebooks that build up from probability and latent variables to sequential VAEs on spike trains, where every step runs and every model is small enough to read.

It grew out of teaching material for our lab, and it assumes you know Python but not necessarily PyTorch, and that you care about neural data specifically — the examples are spikes, behaviour, and latent dynamics rather than MNIST digits.

The notebooks live on GitHub: msenselab/Variational-Neural-Inference. They are meant to be executed, so this page is a map rather than a copy.

The progression

probability
    -> latent variables
    -> mixture models and EM
    -> temporal states and dynamic programming
    -> variational inference
    -> recurrent neural dynamics
    -> nonlinear but locally interpretable dynamics

Each step adds exactly one idea. Mixture models introduce hidden variables. HMMs add temporal structure. VAEs add amortized inference. LFADS adds recurrent latent dynamics for spike trains. gpSLDS adds uncertain nonlinear dynamics assembled from local linear regimes.

The notebooks

#NotebookWhat it covers
00PyTorch PrimerExercise-based PyTorch primer
01PyTorch for NeuroscienceVisual introduction, seminar prerequisite
02Probabilistic ModelingProbability and latent-variable foundations
02bMixture Models and EMMixtures, EM, the ELBO, stochastic EM
03aHMM FoundationsForward-backward, decoding, sampling, Baum-Welch
03From Behaviour to Latent DynamicsGaussian and AR-HMM extended workshop
04Standard VAEThe ELBO and amortized inference
05Variational EMCAVI and coordinate-ascent variational EM
06Sequential VAEs (LFADS)Transparent PyTorch LFADS for spike trains
07Advanced LFADS ReferenceFull JAX workflow, inferring inputs to an integrator RNN
08gpSLDSInterpretable nonlinear dynamics, conceptual capstone

Where to start

If you read four of them, read these: 02 for why latent variables are useful, 04 for the ELBO and amortized inference, 06 for a temporal VAE on real spike trains, and 08 for interpretable nonlinear latent dynamics. Notebooks 00 and 01 are prerequisites rather than lessons. The complete classical inference track is 02 → 02b → 03a.

Notebook 04 is the one I would hand to someone who has read about VAEs and still does not feel they could write one.

Running them

Notebooks 07 and 08 depend on external implementations registered as Git submodules, so downloading single files is not enough for those two:

git clone https://github.com/msenselab/Variational-Neural-Inference.git
cd Variational-Neural-Inference
python scripts/setup_all.py

The setup script initializes the submodules, applies the JAX and NumPy compatibility patches, and installs a CPU-compatible environment for all eleven notebooks. It is idempotent, and --dry-run shows what it would do. For a smaller install, requirements-core.txt covers notebooks 00–06.

A note on honesty about what runs: notebook 05 does its full synthetic CAVI tutorial without any download, and only fetches the Kato dataset if you opt in. Notebook 06 splits 950 training and 50 genuinely held-out test trials. Notebook 08 demonstrates the gpSLDS computational core rather than reproducing the full upstream fitting pipeline. Those boundaries are marked in the notebooks themselves.

Attribution

The tutorial adapts material from Stanford’s STATS 320 and from the original model implementations; references/ATTRIBUTION.md in the repository has the details. The underlying papers are Kingma and Welling (2013) for the VAE, Pandarinath et al. (2018) for LFADS, Vyas et al. (2020) for computation through dynamics, and Hu et al. (2024) for gpSLDS.

← Model Workshop