Variational Autoencoders from First Principles
Latent variable models (LVM) are my favorite thing in machine learning. Their purpose is to model the unobserved processes that generate data. While they are rarely the key to making things work, they provide a pretty universal thought framework for formalizing ML prolems. This makes talking about and checking the correctness of a wide variety of things really easy! Unfortunately, LVMs are really computationally intensive to train. In order to train a LVM, you have to consider all the different ways of generating data in order to find the best one. If that sounds like reinforcement learning, that’s because it is!
In this post we will introduce the math of LVMs, how to train them, and derive variational autoencoders from scratch. This post is longer than usual (for me) and aims to walk you through the problem solving process of training models that reason about the unknown.
Problem setup
We model the following generative process: Sample a latent variable then generate an observation , yielding a model . Our goal is to train this model on a dataset consisting of a single observation (WLOG 1), without observing . This requires maximizing the log marginal likelihood or evidence, where we marginalize over all possible :
Derivation steps
- (1) Total probability
- (2) Model definition
The most convenient approach to optimizing the evidence is to compute it exactly and compute gradients using automatic differentiation. Howeer, computing this sum is intractable if the space of is really large. Optimizing the evidence then becomes a game of approximation. Since we cannot optimize the objective, our goal is to find surrogate losses that are correlated with the evidence and also admit computationally cheap gradient estimators for optimization.
We will explore a series of surrogate objectives, deriving:
- The relationship between the surrogate objective and the evidence
- Gradient estimators to optimize each surrogate objective
Background: Principles of gradient estimation
Gradient estimation attempt 1: Deriving a Monte Carlo estimator
If is parameterized by a transformer, you have to optimize the loss via gradient descent. However, since computing the loss is difficult, we need to find a way to approximate the gradient of the loss. Looking at the loss in equation (2), there’s no really obvious approximation approach yet.
Since there isn’t a clear path forward, let’s just try brute force. As a general rule, approximating an intractable sum requires rewriting that sum as an expectation. We will push the gradient through and see if an expectation falls out.
Derivation steps
- (3) Derivative of log
- (4) Total probability
- (5) Chain rule and product rule
It looks like we are getting somewhere with the first term in the summand, as it seems pretty close to being written as an expectation. There are just some pesky terms lying around that we need to clean up. We can try pushing the derivative further. A classic tool for rewriting gradients as expectations is the log derivative trick:
We can use this to simplify the first term:
Derivation steps
- (6) Log-derivative trick
- (7) Distribute
- (8) Conditional probability
Similarly, we can simplify the second term:
Derivation steps
- (9) Log-derivative trick
- (10) Distribute
- (11) Conditional probability
Combining both terms yields
Exercise: Simplifying the derivation
We can greatly simplify the derivation of the exact gradient estimator by applying the log-derivative trick earlier in the derivation (exercise for the reader, see footnote for a big hint2). This is mostly clear in hindsight, after we see what the expectation looks like. The systematic tool of brute-force derivation is more useful for a first approach. Let’s not pretend things are elegant on the first attempt.
Todo on explaining what is and why its hard to sample from. What happens if we just replace it with something else? In particular, it would be nice to replace with the easy-to-sample .
Aside: Log-scaled rewards
Todo
How can we reason about the quality of the approximation?
Expanding the search space: A surrogate loss
Adding a new term to the evidence to make this more principled. Want this term to drop out to zero when things are finished.
Footnotes
-
You could specialize this to datasets with multiple independent examples as follows. First, by cramming in all examples into a single observation, yielding a single complex observation generated from complex latent . You could then make conditional independence assumptions in the model . ↩
-
Big hint: ↩