It is recommended to read this post as a PDF, available here.
Introduction
Suppose you’re in Vegas, and you’ve had the misfortune of encountering a crooked croupier (let’s call him ). You suspect he is using loaded dice. You’ve been watching him for a while, and you’ve developed a theory for how the dice are loaded (call your theory ). The KL divergence between and is the measure of how surprised you and your wallet would be, on average, if you were betting according to your theory but the dice were behaving according to . This surprise is also a measure of the distance between your and the croupier’s distribution—how far was your guess?
Imagine now that perhaps you rightfully called out this crooked croupier, which led to us getting kicked out of that casino. But of course we must get even. So we find some long-suffering friend of yours. We need a method of long-distance communication. We decide that you’ll blink at him—something like Morse code—from the buffet next door. So we prep him to read blinks, and send him in.
To make the code as efficient as possible, we use information theory and our distribution . If a play appears with probability according to us, the optimal number of bits to encode it is . We’ll note that this assigns fewer bits to more likely things, thus reducing the total amount we need to blink in the direction of our friend.
In the event that our distribution is wrong and the actual distribution is , we should’ve assigned bits to that play instead. So our waste encoded for play would be
If you leave your friend playing for a while at this table, the expected waste is:
This is forward KL divergence—and can be thought of as the wasted bits you need to send over based on the difference between your distribution and the true distribution.
KL Divergence and Entropy
Playing with this equation, we can discover something else quite insightful. First we know that:
Let’s expand that log ratio:
Separate the terms:
By definition,
is the entropy of the true distribution (“entropy of reality”), and
is the cross-entropy of using when samples come from .
Therefore:
KL divergence can also be thought of as the regret, or “surprise tax,” you pay for using the wrong distribution when the true distribution is : it is the gap between the code length you actually incur (cross-entropy ) and the optimal code length you could have achieved if you had known (entropy ).
KL
The code tuned to the true distribution is, in expectation, unbeatable. At best you tie it when ; otherwise you pay the surprise tax. Formally, is the optimal average code length when the world is . is the average code length you get when you insist the world looks like . In expectation, you can’t beat the optimal code, and you only match it when you guessed perfectly. The difference should therefore always be .
This is also equivalent to Gibbs’ Inequality, which we’ll succinctly derive now.
Lemma. For all ,
with equality if and only if .
This is true because is concave. Now apply this to KL. Start with:
Let
so that
Since probability values cannot be negative, , the assumption holds for all values. From for all , we get
Multiply both sides by :
Now sum over all :
since both and are probability distributions and thus each sum to 1. The left-hand side is exactly , so we conclude
with equality if and only if for all , i.e. for all , which means everywhere.
Log Likelihood
Suppose the real world has some unknown distribution , and we build a model with parameters to approximate it. In practice, we fit by maximizing the log-likelihood of the observed data:
This has a close relationship with KL Divergence. Start from the forward KL:
The first term,
depends only on the true distribution , which we do not control. So, as a function of ,
Therefore, when optimizing:
Maximum likelihood training is choosing the model whose predictions make the observed world least surprising on average. Phrased yet another way: among all , we pick the one that wastes the fewest extra bits compared to the (unknowable) true compressor for . The closer is to , the better the model is able to compress its data distribution, and the more it “understands”.
Forward vs Reverse KL Divergence
Where forward KL is , reverse KL is . While their equations look nearly identical, the behavior of a policy iterating under either KL could not be more different.
Forward KL is mode-covering: in a multi-modal distribution, it tries to split the difference and cover as much as possible. Imagine , and , then . Even a tiny bit of probability in , when says “impossible,” makes KL divergence infinite, so forward KL spreads the distribution out.
Reverse KL is mode-seeking: it tends to pick one mode of and match it perfectly. When , there is no penalty. But when and , then . So, reverse KL says: “You can ignore regions where is small, but you absolutely cannot claim something is possible when it is actually impossible.”
Forward KL is used by default in many contexts. Reverse KL is used in generative models like VAEs, where we’d like clear, sharp faces from specific ethnicity/ages, rather than blurry “average” human faces. We also use reverse KL in model distillation, where a large model might say an answer “could be A, B, or C” and you want your small model to model “definitely A” instead of being unable to parse the nuance of A/B/C and being unable to learn at all.
Figure — Evolution of under forward KL (mode-covering) versus reverse KL (mode-seeking) optimization. Left: Initial configuration with starting between two modes of . Right: Final convergence after optimization—forward KL spreads to cover both peaks while reverse KL commits to matching a single mode perfectly.
KL Divergence Estimators
This section is heavily inspired by this blog post. It is slightly lighter on mathematical theory than the source, and puts slightly more effort into motivating the various estimators—all flaws my own.
In RLHF, KL Divergence is used to prevent models from going completely off the rails. We have a fixed reference model and the updating policy . We define as the distribution for position conditioned on all previous tokens up to . The full KL penalty would be:
Computing this exactly would require evaluating all probabilities and for every token in the vocabulary, at every position in the sequence, for every sequence in the batch. With typical values (vocab size = 50,000, sequence length = 2,048, batch size = 32), this would be roughly 3.3 billion probability evaluations per batch—which can be too memory- or computationally-inefficient. So, we need to estimate the value instead.
A good estimator is unbiased (it has the same mean as the original) and preferably has low variance.
A naive estimator would be:
It is unbiased, but it has very high variance. This value can often be negative, even though .
We can sample from , and for each sample we can compute
Any estimator we build has to be some function . So our goal is to pick such that approximates KL well.
Let to keep things clean. When , we have . Ideally, our estimator should have the following properties:
- When , KL is zero. So we want .
- locally matches KL when and are close to each other (often the case in practice). When and are close, small perturbations don’t matter, so we want .
- It has lower per-sample variance than . Here, in order to avoid a dive into Fisher information theory, we assume that the naive estimator has second derivative 1 in the right coordinates. To measure distance on the same scale, we want (read the original blog if you want to dig further in).
Looking at the Taylor expansion of around 0:
Given our constraints, we have , and .
So
and given we want the simplest , we drop the higher-order terms and get
Some nice things fall out of this estimator:
- It’s always positive (like our true KL).
- It measures a distance between and .
- It has lower variance than our naive estimator.
We also have
for small deviations between and . As you will note, this is not unbiased—though the bias is small in practice. We can be quite happy with our estimator.
But we can yet do better! Is there a way to make an unbiased estimator with lower variance? Quoting the original blog: “The general way to lower variance is with a control variate—take and add something that has expectation 0 but is negatively correlated with .”
What do we know that might have expectation zero? Well, we know that . And so is guaranteed to have zero expectation. If we can find a such that
has lower variance, we’ll have a lower-variance, unbiased estimator.
Calculating the optimal is hard, but we can estimate a reasonable value of to be 1 (see the original blog for why). This gives the estimator:
This is an example of a Bregman divergence—the gap between a convex curve and the tangent line drawn from the curve at some point .
Further Readings of Note
For the curious reader who wants to pursue an even deeper understanding, this document lacks coverage on the following:
- The relationship between log-likelihood and KL
- -divergences and Bregman divergences
- Local geometry and Fisher information
I’d welcome any amendments, fixes, or improvements to this document.