Seyed Masoud Hosseini · Overview · Study log · Weekly summaries · Ideas · Search · Transcript · RSS feed

Deep Reinforcement Learning · Lecture 75 of 99 · 17:52

Lecture 18, Variational Inference, Part 3

CS 285: Lecture 18, Variational Inference, Part 3 on YouTube

Study guide

What this lecture covers

This part solves the scaling problem left at the end of part 2: rather than fitting a separate posterior q_i(z) for every data point, it trains a single inference network q_phi(z|x) that outputs a mean and variance for any input x. This "amortized" inference lets variational inference scale to large datasets with a generative network p_theta(x|z) and an inference network trained together.

The lecture then addresses the harder question of how to compute the gradient of the evidence lower bound with respect to the inference network's parameters phi. It shows that a policy-gradient-style estimator works but suffers from high variance, and introduces the reparameterization trick as a lower-variance alternative available specifically because, unlike in reinforcement learning, the "dynamics" here (the decoder network) are fully known and differentiable.

Key ideas

  • Amortized variational inference: replaces a per-data-point posterior with one neural network q_phi(z|x) (an encoder) paired with a generative network p_theta(x|z) (a decoder), so the cost of inference is shared across all data points.
  • Two gradient steps per update: sample z from q_phi(z|x_i), take a gradient step on theta using log p_theta(x_i|z), then take a gradient step on phi to increase the same evidence lower bound.
  • Policy-gradient estimator for phi: the expectation term in the bound has the same form as a policy gradient objective, so its gradient can be estimated by sampling z and weighting grad_phi log q_phi by a reward-like quantity r, but this estimator has high variance.
  • Reparameterization trick: rewrites z as a deterministic function of phi plus an independent noise term, z = mu_phi(x) + epsilon * sigma_phi(x) with epsilon ~ N(0,1), so the gradient can be computed by backpropagating directly through r instead of through log q_phi.
  • Why RL cannot use this trick: policy gradient is needed in RL because the environment dynamics are unknown and non-differentiable; in amortized VI, the decoder is a known, differentiable neural network, so direct differentiation is possible.
  • Closed-form KL term: when the prior p(z) is Gaussian, the entropy and KL-divergence terms in the bound have closed-form expressions in terms of the means and variances, requiring no sampling.
  • Trade-off: the reparameterization trick only applies to continuous latent variables and has lower variance with a single sample; the policy-gradient estimator works for discrete variables too but needs more samples and smaller learning rates.

Before you watch

  • Watch parts 1 and 2 of this lecture, which introduce latent variable models and derive the evidence lower bound this part optimizes.
  • Review the policy gradient lecture from earlier in the course, since the reparameterization trick is explained by contrast with the policy-gradient-style estimator.

Check your understanding

  1. What problem does amortized variational inference solve compared to fitting a separate q_i(z) per data point?
  2. Why does the policy-gradient-style estimator for grad_phi of the evidence lower bound tend to have high variance?
  3. How does the reparameterization trick rewrite z so that phi only affects deterministic quantities, and why does this let you backpropagate directly through r?
  4. Why is the reparameterization trick generally not available for standard reinforcement learning, but is available here?

Vocabulary

amortized inference (noun)
Using one trained network to quickly infer results for any input, instead of solving separately each time.
Amortized inference scales variational inference to large data sets.
inference network (noun)
A neural network trained to estimate a hidden variable from an input.
The inference network outputs a mean and variance for z.
encoder (noun)
A network that compresses an input into a smaller hidden representation.
The encoder maps x to a distribution over z.
decoder (noun)
A network that turns a hidden representation back into an output.
The decoder reconstructs x from z.
policy gradient (noun)
A method that estimates how to change a policy's parameters to increase expected reward.
The gradient looks like a policy gradient estimator here.
high variance (adjective)
Producing results that fluctuate a lot from one sample to another.
The policy-gradient estimator has high variance.
reparameterization trick (noun)
A method that rewrites a random variable as a fixed function plus separate random noise, so gradients can pass through it.
The reparameterization trick gives a lower-variance gradient.
backpropagate (verb)
To compute gradients by passing them backward through a network.
We can backpropagate directly through the reward with this trick.
differentiable (adjective)
Able to have a gradient calculated for it.
The decoder is a known, differentiable network.
noise (noun)
Random values added to a system, often to represent randomness.
Epsilon is an independent noise term drawn from a normal distribution.
trade-off (noun)
A balance between two good things where gaining one costs some of the other.
There is a trade-off between the two gradient estimators.
posterior (noun)
An updated probability distribution over a hidden variable after seeing data.
Amortized inference avoids fitting a separate posterior for every data point.
generative network (noun)
A network that produces new data samples from a hidden representation.
The generative network p_theta(x|z) is trained alongside the encoder.
evidence lower bound (noun)
A quantity that is easier to optimize and is guaranteed to be no higher than the true data likelihood.
Both gradient steps aim to increase the evidence lower bound.
latent variable (noun)
A hidden variable that is not directly observed but explains the data.
z is the latent variable that the inference network estimates.
KL divergence (noun)
A measure of how different one probability distribution is from another.
The KL-divergence term has a closed-form expression when the prior is Gaussian.
closed-form (adjective)
Expressible as an exact formula, without needing to estimate it by sampling.
The entropy term has a closed-form expression here.
entropy (noun)
A measure of how spread out or uncertain a probability distribution is.
The entropy term can be computed exactly for a Gaussian.
gradient step (noun)
One small update to a model's parameters in the direction that improves it.
The method takes one gradient step on theta and one on phi.
scale (verb)
To keep working well as the amount of data or the size of the problem grows.
Amortized inference lets variational inference scale to large datasets.
estimator (noun)
A method or formula that gives an approximate value for a quantity you cannot compute exactly.
The policy-gradient estimator has much higher variance than the reparameterized one.
sample (verb)
To draw a random value from a probability distribution.
We sample z from q_phi(z|x_i) before taking a gradient step.
continuous (adjective)
Able to take any value in a range, not just separate distinct values.
The reparameterization trick only applies to continuous latent variables.
discrete (adjective)
Made up of separate, distinct values rather than a smooth range.
The policy-gradient estimator also works for discrete variables.
learning rate (noun)
The size of each step taken when updating a model's parameters.
The policy-gradient estimator needs smaller learning rates to stay stable.
Gaussian (adjective)
Following the bell-shaped normal probability distribution.
The prior p(z) is assumed to be Gaussian.

Chapters

← Lecture 18, Variational Inference, Part 2 · Lecture 18, Variational Inference, Part 4 →