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 pp). 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 qq). The KL divergence between qq and pp is the measure of how surprised you and your wallet would be, on average, if you were betting according to your theory qq but the dice were behaving according to pp. 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 qq. If a play xx appears with probability q(x)q(x) according to us, the optimal number of bits to encode it is −log⁡2q(x)-\log_2 q(x). 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 pp, we should’ve assigned −log⁡2p(x)-\log_2 p(x) bits to that play instead. So our waste encoded for play xx would be

−log⁡2q(x)−(−log⁡2p(x))=log⁡2p(x)q(x). -\log_2 q(x) - (-\log_2 p(x)) = \log_2\frac{p(x)}{q(x)}.

If you leave your friend playing for a while at this table, the expected waste is:

Ex∼p[log⁡p(x)q(x)]=∑xp(x)⋅log⁡2p(x)q(x).\begin{equation} \mathbb{E}_{x \sim p} \left[ \log\frac{p(x)}{q(x)} \right] = \sum_x p(x) \cdot \log_2 \frac{p(x)}{q(x)}. \end{equation}

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:

KL[p∥q]=∑xp(x)log⁡p(x)q(x).\begin{align} \mathrm{KL}[p\|q] &= \sum_x p(x) \log\frac{p(x)}{q(x)}. \end{align}

Let’s expand that log ratio:

KL[p∥q]=∑xp(x) [log⁡p(x)−log⁡q(x)].\begin{align} \mathrm{KL}[p\|q] &= \sum_x p(x)\,[\log p(x) - \log q(x)]. \end{align}

Separate the terms:

KL[p∥q]=∑xp(x)log⁡p(x)−∑xp(x)log⁡q(x).\begin{align} \mathrm{KL}[p\|q] = \sum_x p(x) \log p(x) - \sum_x p(x) \log q(x). \end{align}

By definition,

H(p)=−∑xp(x)log⁡p(x) H(p) = - \sum_x p(x)\log p(x)

is the entropy of the true distribution (“entropy of reality”), and

Hp(q)=−∑xp(x)log⁡q(x) H_p(q) = - \sum_x p(x)\log q(x)

is the cross-entropy of using qq when samples come from pp.

Therefore:

KL[p∥q]=Hp(q)−H(p)\begin{equation} \boxed{\mathrm{KL}[p\|q] = H_p(q) - H(p)} \end{equation}

KL divergence can also be thought of as the regret, or “surprise tax,” you pay for using the wrong distribution qq when the true distribution is pp: it is the gap between the code length you actually incur (cross-entropy Hp(q)H_p(q)) and the optimal code length you could have achieved if you had known pp (entropy H(p)H(p)).

KL ≥0\geq 0

The code tuned to the true distribution pp is, in expectation, unbeatable. At best you tie it when q=pq = p; otherwise you pay the surprise tax. Formally, H(p)H(p) is the optimal average code length when the world is pp. Hp(q)H_p(q) is the average code length you get when you insist the world looks like qq. In expectation, you can’t beat the optimal code, and you only match it when you guessed perfectly. The difference KL[p∥q]\mathrm{KL}[p\|q] should therefore always be ≥0\ge 0.

This is also equivalent to Gibbs’ Inequality, which we’ll succinctly derive now.

Lemma. For all x>0x > 0,

log⁡x≤x−1, \log x \le x - 1,

with equality if and only if x=1x = 1.

This is true because log⁡x\log x is concave. Now apply this to KL. Start with:

KL[p∥q]=∑xp(x)log⁡p(x)q(x). \mathrm{KL}[p\|q] = \sum_x p(x)\log \frac{p(x)}{q(x)}.

Let

u(x)=q(x)p(x), u(x) = \frac{q(x)}{p(x)},

so that

log⁡p(x)q(x)=−log⁡u(x). \log \frac{p(x)}{q(x)} = -\log u(x).

Since probability values cannot be negative, p(x)≥0,q(x)≥0p(x) \geq 0, q(x) \geq 0, the assumption u(x)>0u(x) > 0 holds for all values. From log⁡u≤u−1\log u \le u - 1 for all u>0u > 0, we get

−log⁡u(x)≥1−u(x). -\log u(x) \ge 1 - u(x).

Multiply both sides by p(x)p(x):

p(x)log⁡p(x)q(x)≥p(x)(1−u(x))=p(x)−q(x). p(x)\log \frac{p(x)}{q(x)} \ge p(x)\bigl(1 - u(x)\bigr) = p(x) - q(x).

Now sum over all xx:

∑xp(x)log⁡p(x)q(x)≥∑x(p(x)−q(x))=∑xp(x)−∑xq(x)=1−1=0 \begin{aligned} \sum_x p(x)\log \frac{p(x)}{q(x)} &\ge \sum_x \bigl(p(x) - q(x)\bigr) &= \sum_x p(x) - \sum_x q(x) \\ &= 1 - 1 = 0 \end{aligned}

since both pp and qq are probability distributions and thus each sum to 1. The left-hand side is exactly KL[p∥q]\mathrm{KL}[p\|q], so we conclude

KL[p∥q]≥0, \mathrm{KL}[p\|q] \ge 0,

with equality if and only if log⁡u(x)=u(x)−1\log u(x) = u(x) - 1 for all xx, i.e. u(x)=1u(x) = 1 for all xx, which means p(x)=q(x)p(x) = q(x) everywhere.

Log Likelihood

Suppose the real world has some unknown distribution p(x)p(x), and we build a model qθ(x)q_\theta(x) with parameters θ\theta to approximate it. In practice, we fit θ\theta by maximizing the log-likelihood of the observed data:

max⁡θ  Ex∼p[log⁡qθ(x)]. \max_\theta \; E_{x \sim p}[\log q_\theta(x)].

This has a close relationship with KL Divergence. Start from the forward KL:

KL[p∥qθ]=∑xp(x)log⁡p(x)qθ(x)=∑xp(x)log⁡p(x)−∑xp(x)log⁡qθ(x). \begin{aligned} \mathrm{KL}[p\|q_\theta] &= \sum_x p(x)\log \frac{p(x)}{q_\theta(x)} \\ &= \sum_x p(x)\log p(x) - \sum_x p(x)\log q_\theta(x). \end{aligned}

The first term,

∑xp(x)log⁡p(x), \sum_x p(x)\log p(x),

depends only on the true distribution pp, which we do not control. So, as a function of θ\theta,

KL[p∥qθ]=constant−Ex∼p[log⁡qθ(x)]. \mathrm{KL}[p\|q_\theta] = \text{constant} - E_{x \sim p}[\log q_\theta(x)].

Therefore, when optimizing:

θ⋆=arg⁡max⁡θEx∼p[log⁡qθ(x)]=arg⁡min⁡θKL[p∥qθ]. \begin{aligned} \theta^\star &= \arg\max_\theta E_{x \sim p}[\log q_\theta(x)] \\ &= \arg\min_\theta \mathrm{KL}[p\|q_\theta]. \end{aligned}

Maximum likelihood training is choosing the model whose predictions make the observed world least surprising on average. Phrased yet another way: among all qθq_\theta, we pick the one that wastes the fewest extra bits compared to the (unknowable) true compressor for pp. The closer qθq_\theta is to pp, 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 KL[reality∥guess]\mathrm{KL}[\text{reality}\|\text{guess}], reverse KL is KL[guess∥reality]\mathrm{KL}[\text{guess}\|\text{reality}]. 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 p(x)=0.01p(x) = 0.01, and q(x)=0.0q(x) = 0.0, then log⁡(p(x)/q(x))=∞\log(p(x)/q(x)) = \infty. Even a tiny bit of probability in pp, when qq 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 pp and match it perfectly. When q(x)=0q(x) = 0, there is no penalty. But when p(x)=0p(x) = 0 and q(x)>0q(x) > 0, then KL=∞\mathrm{KL} = \infty. So, reverse KL says: “You can ignore regions where pp 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 qq under forward KL (mode-covering) versus reverse KL (mode-seeking) optimization. Left: Initial configuration with qq starting between two modes of pp. 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 πref\pi_\text{ref} and the updating policy π\pi. We define π(xt∣x<t)\pi(x_t \mid x_{<t}) as the distribution for position tt conditioned on all previous tokens up to tt. The full KL penalty would be:

KL penalty=KL[π(xt∣x0−t) ∣∣πref(xt∣x0t)]=∑v∈vocabπ(v∣x0−t)log⁡π(v∣x0−t)πref(v∣x0−t) \begin{aligned} \text{KL penalty} &= \mathrm{KL}[\pi(x_t \mid x_{0-t})\ || \pi_{\text{ref}}(x_t \mid x_{0t})] \\ &= \sum_{v \in \text{vocab}} \pi(v \mid x_{0-t}) \log \frac{\pi(v \mid x_{0-t})}{\pi_{\text{ref}}(v \mid x_{0-t})} \end{aligned}

Computing this exactly would require evaluating all probabilities π(v∣x<t)\pi(v \mid x_{<t}) and πref(v∣x<t)\pi_\text{ref}(v \mid x_{<t}) for every token vv in the vocabulary, at every position tt 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:

k1=−log⁡p(x)q(x)=−log⁡r,KL^=E[k1]=E[−log⁡r]. \begin{aligned} k_1 &= - \log \frac{p(x)}{q(x)} = - \log r, \\ \hat{\mathrm{KL}} &= E[k_1] = E[- \log r]. \end{aligned}

It is unbiased, but it has very high variance. This value can often be negative, even though KL≥0\mathrm{KL} \geq 0.

We can sample from qq, and for each sample xx we can compute

log⁡r(x)=log⁡p(x)q(x). \log r(x) = \log \frac{p(x)}{q(x)}.

Any estimator we build has to be some function g(log⁡r)g(\log r). So our goal is to pick gg such that Eq[g(log⁡r)]E_q[g(\log r)] approximates KL well.

Let t=log⁡rt = \log r to keep things clean. When p=qp = q, we have r=1r = 1. Ideally, our estimator should have the following properties:

  1. When p=qp = q, KL is zero. So we want g(0)=0g(0) = 0.
  2. g(⋅)g(\cdot) locally matches KL when pp and qq are close to each other (often the case in practice). When pp and qq are close, small perturbations don’t matter, so we want g′(0)=0g'(0) = 0.
  3. It has lower per-sample variance than −log⁡r- \log r. Here, in order to avoid a dive into Fisher information theory, we assume that the naive estimator −log⁡r- \log r has second derivative 1 in the right coordinates. To measure distance on the same scale, we want g′′(0)=1g''(0) = 1 (read the original blog if you want to dig further in).

Looking at the Taylor expansion of gg around 0:

g(t)=a0+a1t+12a2t2+O(t3).\begin{align*} g(t) = a_0 + a_1 t + \tfrac{1}{2} a_2 t^2 + O(t^3). \end{align*}

Given our constraints, we have a0=g(0)=0a_0 = g(0) = 0, a1=g′(0)=0a_1 = g'(0) = 0 and a2=g′′(0)=1a_2 = g''(0) = 1.

So

g(t)=12t2+O(t3), g(t) = \tfrac{1}{2} t^2 + O(t^3),

and given we want the simplest gg, we drop the higher-order terms and get

g(log⁡r)=12(log⁡r)2. g(\log r) = \frac{1}{2} (\log r)^2.

Some nice things fall out of this estimator:

  • It’s always positive (like our true KL).
  • It measures a distance between pp and qq.
  • It has lower variance than our naive estimator.

We also have

E[K2]=KL(q∥p)+O(δ3), E[K_2] = \mathrm{KL}(q\|p) + O(\delta^3),

for small deviations δ\delta between pp and qq. As you will note, this is not unbiased—though the bias is small in practice. We can be quite happy with our k2k_2 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 k1k_1 and add something that has expectation 0 but is negatively correlated with k1k_1.”

What do we know that might have expectation zero? Well, we know that E[r]=1E[r] = 1. And so (r−1)(r - 1) is guaranteed to have zero expectation. If we can find a λ\lambda such that

−log⁡r+λ(1−r) -\log r + \lambda (1 - r)

has lower variance, we’ll have a lower-variance, unbiased estimator.

Calculating the optimal λ\lambda is hard, but we can estimate a reasonable value of λ\lambda to be 1 (see the original blog for why). This gives the k3k_3 estimator:

k3=(r−1)−log⁡r. k_3 = (r - 1) - \log r.

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 xx.

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
  • ff-divergences and Bregman divergences
  • Local geometry and Fisher information

I’d welcome any amendments, fixes, or improvements to this document.