DEV Community

Dinesh Kumar Sarangapani
Dinesh Kumar Sarangapani

Posted on Originally published at dineshkumars.dev

Section 3a: f-divergences

I am working through Mathematical Foundations of Generative AI, Prof. Prathosh AP's public lectures. The playlist is the spine. These notes are my deep dive on each section: a visual when the picture is the point, a detour when a prerequisite is doing real work, and the formula in my own words.

Previously: Section 2: The general principle of generative models.

Section 3a: f-divergences

Section 2 left a slot for a divergence. This note fills it with a family. One convex function $f$ picks the score. Section 3b, not written yet, is the convex conjugate: the rewrite that estimates $D_f$ from samples, which is the bridge into GANs.

The statement I am carrying:

Given two distributions with densities $P_x$ and $P_\theta$,

$$
D_f(P_x \,|\, P_\theta) = \int_{\mathcal{X}} P_\theta(x)\, f!\left(\frac{P_x(x)}{P_\theta(x)}\right) dx
$$

where $f: \mathbb{R}+ \to \mathbb{R}$ is convex and lower-semicontinuous, with $f(1) = 0$. Then $D_f \ge 0$, and $D_f = 0$ if and only if $P_x = P\theta$.

Intuition

At each point $x$ I compare how much probability the data puts there with how much the model puts there. The ratio $P_x / P_\theta$ is 1 where they agree and something else where they do not. $f$ turns that local mismatch into a penalty. The integral averages the penalties. A different $f$ punishes a different kind of mismatch, so one template gives a family.

The load test, by endpoint

I compare synthetic traffic with the real access log, endpoint by endpoint.

$$
\text{ratio} = \frac{\text{real share of traffic}}{\text{synthetic share of traffic}}
$$

A ratio of 1 means the generator gets that endpoint right. A ratio of 3 means real users hit it three times as often as the generator does. A ratio of 0.2 means the generator hammers an endpoint real users barely touch.

The single score depends on which failure I care about.

  • Missing real traffic patterns. The penalty grows when the ratio is large. That is forward KL.
  • Fake traffic that real users would never send. The penalty grows when the ratio is near 0. That is reverse KL.
  • A balanced score that cannot run off to infinity. That is Jensen–Shannon or total variation.

$f$ is the penalty. The integral is the average over endpoints.

Visual

$P_x$ is a standard bell centred at 0. $P_\theta$ is the same bell shifted to $\mu$. Pick an $f$. The top curve is the penalty, and it passes through $(1, 0)$. The coral curve is the integrand $P_\theta(x)\, f(P_x(x)/P_\theta(x))$. The four numbers are the integrals.

Three things I want in front of me while I move $\mu$:

  1. At $\mu = 0$ every divergence is 0. Matching distributions means zero penalty, for every $f$ with $f(1) = 0$.
  2. With forward KL, the coral curve dips below zero in places. A single point can contribute a negative amount. The total never does. The proof below is why.
  3. Jensen–Shannon and total variation level off as $\mu$ grows. Both are bounded. The two KLs keep climbing. That saturation is the later reason a GAN, which aims at Jensen–Shannon, can stop learning when the data and the model do not overlap.

Detours

Density ratio. $u(x) = P_x(x) / P_\theta(x)$ asks, at $x$, how many times more likely real data is here than model data.

  • $u = 1$: agreement at this point.
  • $u > 1$: the model under-produces here.
  • $u < 1$: the model over-produces here.
  • $u = 0$: the model puts mass where the data never goes.

If $P_x = P_\theta$ everywhere, then $u(x) = 1$ everywhere. $f$ only ever sees this ratio.

Convexity. A function is convex when it is bowl-shaped: the straight line between any two points on its graph lies on or above the curve. For $\lambda \in [0, 1]$,

$$
\lambda f(u_1) + (1-\lambda) f(u_2) \;\ge\; f\big(\lambda u_1 + (1-\lambda) u_2\big)
$$

Averaging the outputs is at least as big as the output of the average. Check with $f(u) = u^2$, $u_1 = 0$, $u_2 = 2$, $\lambda = \tfrac{1}{2}$. The average of the outputs is $\tfrac{1}{2}(0 + 4) = 2$. The output of the average is $f(1) = 1$. And $2 \ge 1$.

Convexity is what forces $D_f \ge 0$.

Expectation as an integral. For a continuous density $p$, the expected value of $h(x)$ is a probability-weighted average:

$$
\mathbb{E}{x \sim p}[h(x)] = \int{\mathcal{X}} p(x)\, h(x)\, dx
$$

Jensen's inequality. For convex $f$ and a random variable $U$,

$$
\mathbb{E}[f(U)] \;\ge\; f(\mathbb{E}[U])
$$

This is the two-point definition extended to any average. On a bowl, the average of points on the surface lands above the value at the average, never below.

The formula

$$
D_f(P_x \,|\, P_\theta) = \int_{\mathcal{X}} P_\theta(x)\, f!\left(\frac{P_x(x)}{P_\theta(x)}\right) dx = \mathbb{E}{x \sim P\theta}!\left[f!\left(\frac{P_x(x)}{P_\theta(x)}\right)\right]
$$

Read aloud: the average penalty, over points drawn from the model, of how mismatched the two densities are at each point.

Symbol What it is Type / shape Role
$P_x(x)$ true density at $x$ scalar $\ge 0$ the data's weight here
$P_\theta(x)$ model density at $x$ scalar $\ge 0$ the model's weight, and the weight in the average
$u = P_x/P_\theta$ density ratio scalar $\ge 0$ local mismatch
$f$ generator of the divergence convex, $f(1) = 0$ turns mismatch into a penalty
$\int \cdot\, dx$ integral over $\mathcal{X}$ operation adds the penalties
$D_f$ the divergence scalar $\ge 0$ the total score

Each condition on $f$ earns its place:

  • $f(1) = 0$. If the distributions match, $u = 1$ everywhere, so $D_f = \int P_\theta \cdot 0\, dx = 0$.
  • Convex. This is the $D_f \ge 0$ proof below.
  • Lower-semicontinuous. No sudden downward jumps. Section 3b needs it so the conjugate is well behaved.

The proof that $D_f \ge 0$

The lectures state the property. This is the short proof.

$$
\begin{aligned}
D_f(P_x \,|\, P_\theta) &= \mathbb{E}{x \sim P\theta}!\left[f!\left(\tfrac{P_x(x)}{P_\theta(x)}\right)\right] \
&\ge f!\left(\mathbb{E}{x \sim P\theta}!\left[\tfrac{P_x(x)}{P_\theta(x)}\right]\right) \
&= f!\left(\int P_\theta(x)\, \tfrac{P_x(x)}{P_\theta(x)}\, dx\right) \
&= f!\left(\int P_x(x)\, dx\right) \
&= f(1) = 0
\end{aligned}
$$

The first line is the expectation form. The second is Jensen. The third writes the expectation as an integral. $P_\theta$ cancels. A density integrates to 1, and $f(1) = 0$.

In words: the average ratio, weighted by the model, is exactly 1. A bowl-shaped penalty averaged around 1 cannot fall below its value at 1, which is 0. That is why the coral curve can go negative locally while the total stays non-negative.

One template, four choices of $f$

Name $f(u)$ Resulting divergence Behaviour
Forward KL $u \log u$ $\displaystyle\int P_x \log\frac{P_x}{P_\theta}\, dx = D_{\mathrm{KL}}(P_x \,|\, P_\theta)$ Punishes the model for missing real data. Mode-covering. This is the blurry blob in Section 2, and it is maximum likelihood.
Reverse KL $-\log u$ $\displaystyle\int P_\theta \log\frac{P_\theta}{P_x}\, dx = D_{\mathrm{KL}}(P_\theta \,|\, P_x)$ Punishes the model for generating where the data is not. Mode-seeking. It sits on one hump.
Jensen–Shannon $\tfrac{1}{2}\big[u \log u - (u+1)\log\tfrac{u+1}{2}\big]$ symmetric, bounded by $\log 2$ What the original GAN minimises.
Total variation $\tfrac{1}{2}\lvert u - 1 \rvert$ $\tfrac{1}{2}\int \lvert P_x - P_\theta \rvert\, dx$ Symmetric, bounded by 1.

Forward KL is the template with $f(u) = u \log u$, and the $P_\theta$ cancels:

$$
\int P_\theta \cdot \frac{P_x}{P_\theta} \log\frac{P_x}{P_\theta}\, dx = \int P_x \log\frac{P_x}{P_\theta}\, dx
$$

Reverse KL is the same template with $f(u) = -\log u$. Substituting it recovers $\int P_\theta \log(P_\theta / P_x)\, dx$. The Jensen–Shannon $f$ above satisfies $f(1) = 0$. A GAN writeup often uses $f(u) = u \log u - (u+1)\log(u+1)$, which differs from this one by a constant shift and is why that loss is described as similar to Jensen–Shannon rather than identical. Total variation is $\tfrac{1}{2}\lvert u - 1 \rvert$, which is the choice that turns the template into $\tfrac{1}{2}\int \lvert P_x - P_\theta \rvert$.

A worked example small enough to do by hand

Two outcomes. $P_x = (0.5, 0.5)$, a fair coin. $P_\theta = (0.8, 0.2)$, a biased model.

The ratios are $u_1 = 0.5 / 0.8 = 0.625$ (over-produced) and $u_2 = 0.5 / 0.2 = 2.5$ (under-produced). The template on a finite set is $D_f = \sum_i P_\theta(i)\, f(u_i)$.

Divergence Computation Value
Forward KL $0.8(0.625 \ln 0.625) + 0.2(2.5 \ln 2.5) = 0.8(-0.294) + 0.2(2.291)$ $\approx 0.223$
Reverse KL $0.8(-\ln 0.625) + 0.2(-\ln 2.5) = 0.8(0.470) + 0.2(-0.916)$ $\approx 0.193$
Jensen–Shannon $0.8\, f(0.625) + 0.2\, f(2.5) = 0.8(0.0218) + 0.2(0.1660)$ $\approx 0.051$
Total variation $0.8 \cdot \tfrac{1}{2}(0.375) + 0.2 \cdot \tfrac{1}{2}(1.5) = 0.15 + 0.15$ $0.300$

All four are positive, and they disagree on the size of the gap. Forward KL is not reverse KL, so the score is not symmetric. In the forward-KL row, outcome 1 contributes about $-0.235$ and outcome 2 about $+0.458$. Negative pieces, positive total. That is Jensen. If $P_\theta = (0.5, 0.5)$, every ratio is 1 and every row is 0.

The catch, which is Section 3b

Computing $D_f$ needs the values $P_x(x)$ and $P_\theta(x)$.

  • $P_x$ is unknown. Section 1 only gave me samples.
  • $P_\theta$ is implicit. Section 2 lets me sample $g_\theta(z)$, not evaluate a density.

So I cannot form the ratio $u(x)$. What I can form are averages over samples, by the law of large numbers. Section 3b has to rewrite $D_f$ as expectations under $P_x$ and under $P_\theta$, with no density values in the formula. The convex conjugate is the tool. The result is the GAN objective.

Where this sits in the lectures

  • This note answers "which divergence?" with a menu, not a single choice.
  • A GAN picks the Jensen–Shannon-like $f$.
  • VAEs, diffusion, and autoregressive models minimise forward KL, which is maximum likelihood.
  • Bounded scores such as Jensen–Shannon saturate when $P_x$ and $P_\theta$ do not overlap. Wasserstein, which is not an f-divergence, is the later answer to that.
  • The KL term that keeps an aligned model close to a reference model, $D_{\mathrm{KL}}(\pi_\theta \,|\, \pi_{\mathrm{ref}})$, is this same family.

Questions I want to be able to answer:

  1. Why does $f(1) = 0$ have to hold? What breaks if $f(1) = 5$?
  2. In the worked example, outcome 1 contributes a negative amount to forward KL. Why does that not break $D_f \ge 0$?
  3. Why can I not plug the dataset into the $D_f$ formula and compute it directly?

Top comments (0)