Skip to content

Entropy, cross-entropy, KL, JS, and the k3 estimator

ModuleM11.1 · build · Python · Pass 2 · 2 to 3 h
You buildpython/tinyllm/info/entropy.py: entropy, cross_entropy, kl, kl_from_logprobs, js, kl_k3, entropy_from_logits
Contractcourse/contracts/py/tinyllm/info/entropy.pyi
Testscourse/tests/M11.1/ (what they check: section 4)
NeedsM09.2 stable numerics (log_softmax, or --ref-deps)
Used byM08.3 differentiates this cross-entropy · L0.3 the training loss · later: L8.1 per-step entropy logging, L8.6 acceptance rate as 1−TV1 - \mathrm{TV}, L12.3 the KL penalty with kl_k3, M11.2 perplexity · later: L12.2, L12.4, M11.4
MilestoneMS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains)
Optional depthCover and Thomas, Elements of Information Theory (2nd ed.), ch. 2; MacKay, Information Theory, Inference, and Learning Algorithms, ch. 2 and 4; Schulman, “Approximating KL Divergence” (2020 blog post)
  • Entropy H(p)=−∑ipilog⁡piH(p) = -\sum_i p_i \log p_i is the average surprise of a distribution, between 0 (certain) and log⁡n\log n (uniform) (test_entropy_bounds).
  • Cross-entropy splits as H(p,q)=H(p)+KL(p ∥ q)H(p, q) = H(p) + \mathrm{KL}(p \,\Vert\, q): the training loss is the data’s own entropy, which no model can beat, plus the model’s excess (test_cross_entropy_decomposes).
  • KL≥0\mathrm{KL} \ge 0 with equality only at p=qp = q (Gibbs), and it is not symmetric; Jensen-Shannon is symmetric and bounded by ln⁡2\ln 2 (test_gibbs_kl_nonnegative, test_kl_is_not_symmetric, test_js_properties).
  • From model outputs, compute KL with log-probabilities, never with p/qp / q; the k3 estimator er−r−1e^r - r - 1 estimates KL from samples without bias and is never negative (test_kl_from_logprobs_extreme_logits, test_kl_k3_mean_is_kl).
Terminal window
ol start M11.1 # stubs entropy.py into your repo, contract alongside
ol tests M11.1 # read the test catalog first: rung R0, you write no tests here
ol check M11.1 # exit code is the verdict
ol check M11.1 --ref-deps # only if your M09.2 is not passing yet
ol diff M11.1 # after passing: your code against the reference

The tracer bigram (L0.0) already reports a negative log-likelihood, and train bigram prints it, but nothing in your system says what that number is measured against or what it can reach. In this pass you train the same bigram by gradient descent (L0.5) with a cross-entropy loss (L0.3), and the loss needs a meaning: how far above the floor are we, and what is the floor? Later the sampler (L8.1) logs the entropy of every step so a collapsing temperature is visible in your traces, speculative decoding (L8.6) reads its acceptance rate off the distance between two distributions, and post-training (L12.3) penalizes the policy for drifting from the reference model with a sampled KL. Each needs the same few quantities with the same conventions at zero, in nats, stable on the log-probabilities your models actually produce. And M08.3 needs a forward cross-entropy to check its gradient against.

SymbolMeaningType / shape
nnnumber of outcomes (vocabulary size for a language model)int
p,qp, qprobability distributions over the nn outcomes: pi≥0p_i \ge 0, ∑ipi=1\sum_i p_i = 1float64[n]
log⁡\logthe natural logarithm; results are in nats (divide by ln⁡2\ln 2 for bits)
H(p)H(p)entropy of ppscalar
H(p,q)H(p, q)cross-entropy of qq relative to ppscalar
KL(p ∥ q)\mathrm{KL}(p \,\Vert\, q)Kullback-Leibler divergence from qq to ppscalar, ≥0\ge 0
m=(p+q)/2m = (p + q)/2the mixture of pp and qqfloat64[n]
JS(p,q)\mathrm{JS}(p, q)Jensen-Shannon divergencescalar in [0,ln⁡2][0, \ln 2]
π,πref\pi, \pi_{\mathrm{ref}}a policy and a reference model (L12.3)distributions
x∼πx \sim \pian outcome sampled from π\piindex
r=log⁡πref(x)−log⁡π(x)r = \log \pi_{\mathrm{ref}}(x) - \log \pi(x)log ratio at a samplescalar
k3=er−r−1k_3 = e^r - r - 1the k3 estimator of KL(π ∥ πref)\mathrm{KL}(\pi \,\Vert\, \pi_{\mathrm{ref}})scalar, ≥0\ge 0
zzlogits, with softmax(z)\mathrm{softmax}(z) a distribution (M09.2)float64[n]

Surprise. An outcome of probability pip_i carries surprise −log⁡pi-\log p_i: a certain outcome (pi=1p_i = 1) carries none, a rare one a lot, and the surprise of two independent outcomes adds, because −log⁡(pipj)=−log⁡pi−log⁡pj-\log(p_i p_j) = -\log p_i - \log p_j. With log⁡2\log_2 the unit is the bit: −log⁡2pi-\log_2 p_i is the length of the best code word for outcome ii.

Entropy is the average surprise.

H(p)=−∑ipilog⁡pi.H(p) = -\sum_i p_i \log p_i.

A term with pi=0p_i = 0 is defined as 0, because xlog⁡x→0x \log x \to 0 as x→0x \to 0: an outcome that never happens costs nothing. H(p)≥0H(p) \ge 0 since every surprise is ≥0\ge 0, and H(p)≤log⁡nH(p) \le \log n, with equality exactly for the uniform distribution: nothing is more unpredictable than equal chances.

Cross-entropy measures a model against data. If the data come from pp and you predict with qq, your average surprise is

H(p,q)=−∑ipilog⁡qi.H(p, q) = -\sum_i p_i \log q_i.

This is the language-model loss. With pp the one-hot distribution of the observed next token tt, it is −log⁡qt-\log q_t, and averaging over positions gives the negative log-likelihood your bigram already prints. A term with pi>0p_i > 0 and qi=0q_i = 0 is +∞+\infty: the model said impossible, and it happened.

KL is the excess, and it is never negative. Subtract the entropy:

KL(p ∥ q)=H(p,q)−H(p)=∑ipilog⁡piqi.\mathrm{KL}(p \,\Vert\, q) = H(p, q) - H(p) = \sum_i p_i \log \frac{p_i}{q_i}.

Proof that it is ≥0\ge 0 (Gibbs’ inequality): log⁡y≤y−1\log y \le y - 1 for y>0y > 0, with equality only at y=1y = 1. Over the outcomes with pi>0p_i > 0,

−KL(p ∥ q)=∑ipilog⁡qipi≤∑ipi(qipi−1)=∑i:pi>0qi−1≤0.-\mathrm{KL}(p \,\Vert\, q) = \sum_i p_i \log \frac{q_i}{p_i} \le \sum_i p_i \left(\frac{q_i}{p_i} - 1\right) = \sum_{i: p_i > 0} q_i - 1 \le 0.

Equality needs qi=piq_i = p_i everywhere. So minimizing the cross-entropy over qq minimizes KL, and the loss floor is H(p)H(p): the entropy of the data itself. Training a language model is pushing KL(data ∥ model)\mathrm{KL}(\text{data} \,\Vert\, \text{model}) toward 0.

KL has a direction. KL(p ∥ q)\mathrm{KL}(p \,\Vert\, q) weighs the log ratio by pp, so it punishes qq for being small where pp is large (it is “mass covering”), and is infinite if qq misses any outcome of pp. KL(q ∥ p)\mathrm{KL}(q \,\Vert\, p) punishes the opposite and prefers a qq that sits on one mode of pp. They are different numbers (section 4 has one example), and L12 picks one on purpose.

KL from log-probabilities. Models output log-probabilities (ℓ=logsoftmax(z)\ell = \mathrm{logsoftmax}(z)), which are finite even where the probability underflows. Write

KL(p ∥ q)=∑ieℓip(ℓip−ℓiq),\mathrm{KL}(p \,\Vert\, q) = \sum_i e^{\ell^p_i} \left(\ell^p_i - \ell^q_i\right),

a difference of moderate numbers, and never form eℓip/eℓiqe^{\ell^p_i} / e^{\ell^q_i}, which is 0/00/0 or ∞/∞\infty/\infty when the logits are large. Two edge cases: ℓip=−∞\ell^p_i = -\infty means pi=0p_i = 0 and contributes 0 (even against ℓiq=−∞\ell^q_i = -\infty, where the difference would be NaN); a finite ℓip\ell^p_i against ℓiq=−∞\ell^q_i = -\infty is +∞+\infty, even when eℓipe^{\ell^p_i} underflows to 0 (where 0⋅∞0 \cdot \infty would be NaN).

Jensen-Shannon is a symmetric, bounded cousin. Compare both with their average m=(p+q)/2m = (p + q)/2:

JS(p,q)=12KL(p ∥ m)+12KL(q ∥ m).\mathrm{JS}(p, q) = \tfrac12 \mathrm{KL}(p \,\Vert\, m) + \tfrac12 \mathrm{KL}(q \,\Vert\, m).

mm is positive wherever pp or qq is, so JS is always finite, and since mi≥pi/2m_i \ge p_i / 2, every log⁡(pi/mi)≤log⁡2\log(p_i / m_i) \le \log 2: JS≤ln⁡2\mathrm{JS} \le \ln 2, reached by distributions with disjoint supports.

Estimating KL from samples: k3. When nn is a vocabulary and you only have the log-probabilities of the tokens you sampled, you estimate KL(π ∥ πref)=Ex∼π[log⁡π(x)−log⁡πref(x)]=E[−r]\mathrm{KL}(\pi \,\Vert\, \pi_{\mathrm{ref}}) = \mathbb{E}_{x \sim \pi}[\log \pi(x) - \log \pi_{\mathrm{ref}}(x)] = \mathbb{E}[-r] by averaging over samples. The plain estimator −r-r is unbiased but negative for many samples. Schulman’s k3 adds a term with mean zero:

Ex∼π[er]=∑xπ(x)πref(x)π(x)=∑xπref(x)=1,\mathbb{E}_{x \sim \pi}\left[e^{r}\right] = \sum_x \pi(x) \frac{\pi_{\mathrm{ref}}(x)}{\pi(x)} = \sum_x \pi_{\mathrm{ref}}(x) = 1,

so k3=(er−1)−rk_3 = (e^r - 1) - r has the same mean, KL(π ∥ πref)\mathrm{KL}(\pi \,\Vert\, \pi_{\mathrm{ref}}) (the sum runs over π\pi‘s support, which is everything for softmax outputs). And every sample is ≥0\ge 0, because the exponential lies above its tangent at 0: er≥1+re^r \ge 1 + r.

k3 for small rr. Near r=0r = 0, k3=r2/2+r3/6+…k_3 = r^2/2 + r^3/6 + \dots is tiny while er≈1e^r \approx 1. Computing er−r−1e^r - r - 1 subtracts numbers near 1 and keeps only the rounding error of ere^r, about 10−1610^{-16}, against a true value of 5×10−135 \times 10^{-13} at r=10−6r = 10^{-6}. np.expm1(r) computes er−1e^r - 1 directly to full relative precision, so expm1(r)−r\mathrm{expm1}(r) - r keeps the answer. This is exactly the regime at the start of post-training, when the policy has barely moved.

Entropy from logits. With ℓ=logsoftmax(z)\ell = \mathrm{logsoftmax}(z) from M09.2, H=−∑ieℓiℓiH = -\sum_i e^{\ell_i} \ell_i, with masked entries (ℓi=−∞\ell_i = -\infty) contributing 0. Computing p=softmax(z)p = \mathrm{softmax}(z) and then −∑plog⁡p-\sum p \log p meets 0⋅log⁡0=0⋅(−∞)=NaN0 \cdot \log 0 = 0 \cdot (-\infty) = \mathrm{NaN} as soon as one probability underflows.

Take n=3n = 3, p=[1/2,1/4,1/4]p = [1/2, 1/4, 1/4], q=[1/4,1/4,1/2]q = [1/4, 1/4, 1/2]. Every log is a multiple of ln⁡2=0.693147\ln 2 = 0.693147 except one.

Entropy. H(p)=12ln⁡2+14ln⁡4+14ln⁡4=(12+12+12)ln⁡2=1.5ln⁡2=1.039721H(p) = \tfrac12 \ln 2 + \tfrac14 \ln 4 + \tfrac14 \ln 4 = (\tfrac12 + \tfrac12 + \tfrac12) \ln 2 = 1.5 \ln 2 = 1.039721 nats, which is 1.5 bits.

Cross-entropy. H(p,q)=12ln⁡4+14ln⁡4+14ln⁡2=(1+12+14)ln⁡2=1.75ln⁡2=1.213008H(p, q) = \tfrac12 \ln 4 + \tfrac14 \ln 4 + \tfrac14 \ln 2 = (1 + \tfrac12 + \tfrac14) \ln 2 = 1.75 \ln 2 = 1.213008.

KL. Directly: 12ln⁡1/21/4+14ln⁡1+14ln⁡1/41/2=12ln⁡2−14ln⁡2=0.25ln⁡2=0.173287\tfrac12 \ln \frac{1/2}{1/4} + \tfrac14 \ln 1 + \tfrac14 \ln \frac{1/4}{1/2} = \tfrac12 \ln 2 - \tfrac14 \ln 2 = 0.25 \ln 2 = 0.173287. Check: H(p,q)−H(p)=1.75ln⁡2−1.5ln⁡2H(p, q) - H(p) = 1.75 \ln 2 - 1.5 \ln 2. Same number.

Jensen-Shannon. m=[3/8,1/4,3/8]m = [3/8, 1/4, 3/8].

log⁡(pi/mi)\log(p_i / m_i)pilog⁡(pi/mi)p_i \log(p_i/m_i)
i=1i = 1ln⁡(4/3)=0.287682\ln(4/3) = 0.2876820.1438410.143841
i=2i = 2ln⁡1=0\ln 1 = 00
i=3i = 3ln⁡(2/3)=−0.405465\ln(2/3) = -0.405465−0.101366-0.101366

So KL(p ∥ m)=0.042475\mathrm{KL}(p \,\Vert\, m) = 0.042475; by symmetry of this example KL(q ∥ m)\mathrm{KL}(q \,\Vert\, m) is the same, and JS=0.042475=1.25ln⁡2−0.75ln⁡3\mathrm{JS} = 0.042475 = 1.25 \ln 2 - 0.75 \ln 3, well under ln⁡2\ln 2.

k3. Let π=p\pi = p and πref=q\pi_{\mathrm{ref}} = q. The log ratios r=log⁡(qx/px)r = \log(q_x / p_x) are [−ln⁡2,0,ln⁡2][-\ln 2, 0, \ln 2]:

xxπ(x)\pi(x)rrk3=er−r−1k_3 = e^r - r - 1
11/2−0.693147-0.6931470.5+0.693147−1=0.1931470.5 + 0.693147 - 1 = 0.193147
21/400
31/40.6931470.6931472−0.693147−1=0.3068532 - 0.693147 - 1 = 0.306853

The mean under π\pi: 12⋅0.193147+14⋅0.306853=0.096574+0.076713=0.173287=0.25ln⁡2\tfrac12 \cdot 0.193147 + \tfrac14 \cdot 0.306853 = 0.096574 + 0.076713 = 0.173287 = 0.25 \ln 2. Exactly KL(p ∥ q)\mathrm{KL}(p \,\Vert\, q), as promised, and every row is nonnegative.

Entropy from logits. z=[ln⁡2,0,0]z = [\ln 2, 0, 0] gives ez=[2,1,1]e^z = [2, 1, 1], softmax [1/2,1/4,1/4]=p[1/2, 1/4, 1/4] = p, so the entropy is again 1.5ln⁡21.5 \ln 2.

These numbers are the first case in section 4, test_hand_example.

# python/tinyllm/info/entropy.py (natural logs: nats)
def entropy(p: ArrayLike, axis: int = -1) -> NDArray
def cross_entropy(p: ArrayLike, q: ArrayLike, axis: int = -1) -> NDArray
def kl(p: ArrayLike, q: ArrayLike, axis: int = -1) -> NDArray
def kl_from_logprobs(logp: ArrayLike, logq: ArrayLike, axis: int = -1) -> NDArray
def js(p: ArrayLike, q: ArrayLike) -> NDArray # last axis
def kl_k3(logp_ref: ArrayLike, logp: ArrayLike) -> NDArray # elementwise, no reduction
def entropy_from_logits(z: ArrayLike, axis: int = -1) -> NDArray

Each reduces over axis and drops it, like np.sum. Probabilities are not renormalized, and a negative one is a ValueError. Note the argument order of kl_k3: the reference model first, as in r=log⁡πref−log⁡πr = \log \pi_{\mathrm{ref}} - \log \pi. Compute the zero conventions with np.where on the terms; adding a small epsilon inside the logs biases every answer and turns a true +∞+\infty into a large finite number.

TestKINDChecksWhy it matters downstream
test_hand_exampleunitevery section 3 number, in natsyou and the test agree on the definitions
test_matches_scipy_goldengoldenscipy’s entropy, KL, and JS on dense, sparse, peaked, and column-wise inputsan independent implementation agrees
test_gibbs_kl_nonnegativepropertyKL≥0\mathrm{KL} \ge 0 on 200 random pairs, exactly 0 at p=qp = qa KL penalty never rewards drift (L12.3)
test_cross_entropy_decomposespropertyH(p,q)=H(p)+KL(p ∥ q)H(p, q) = H(p) + \mathrm{KL}(p \,\Vert\, q)reading the loss: floor plus excess
test_kl_is_not_symmetricunit[0.9,0.1][0.9, 0.1] vs uniform: 0.36806 one way, 0.51083 the otherforward and reverse KL are different objectives
test_zero_probability_conventionsboundary0log⁡0=00 \log 0 = 0, plog⁡(p/0)=+∞p \log(p/0) = +\infty, no NaNsparse distributions, masked tokens
test_rejects_negative_probabilitiesboundarya negative entry is a ValueError everywherea logit passed as a probability fails loudly
test_entropy_boundsproperty0 for one-hot, log⁡n\log n for uniform, in between otherwisereading L8.1’s entropy log
test_axis_and_batch_shapesunit[B, V] reduces to [B]; axis=0 reduces columnsbatched losses
test_js_propertiespropertysymmetric, in [0,ln⁡2][0, \ln 2], ln⁡2\ln 2 for disjoint supportsa bounded comparison of two models
test_kl_from_logprobs_matches_kldifferentialequals kl on the same distributions, −∞-\infty entries includedthe form models use
test_kl_from_logprobs_extreme_logitsboundarylog-probabilities from logits near 10410^4 give a finite, correct KLlate training, low temperature
test_kl_from_logprobs_infinite_casesboundary−∞-\infty against −∞-\infty is 0; finite (even −800-800) against −∞-\infty is +∞+\inftyno NaN from ∞−∞\infty - \infty or 0⋅∞0 \cdot \infty
test_kl_k3_is_nonnegativepropertyk3≥0k_3 \ge 0 on 1000 samples, 0 at r=0r = 0, e−2e - 2 at r=1r = 1per-sample penalties are never negative
test_kl_k3_small_r_precisionboundaryrelative error below 10−910^{-9} at ∣r∣≤3×10−5\lvert r\rvert \le 3 \times 10^{-5}the start of L12.3, when the policy barely moved
test_kl_k3_mean_is_klstatisticalexact enumeration equals KL; 20000 samples within 4 standard errorsthe estimator is unbiased
test_entropy_from_logitsboundaryequals H(softmax(z))H(\mathrm{softmax}(z)), finite at 10410^4, masks count as 0L8.1 logs it from logits
PitfallSymptomCaught by
1. p * np.log(p) with p=0p = 0, or an epsilon inside the logsNaN entropy for sparse rows; a finite KL where it is infinitetest_zero_probability_conventions (mutants s01, s02)
2. forming p/qp/q from exponentiated log-probabilities0/0 and inf/inf once logits are largetest_kl_from_logprobs_extreme_logits (mutant s05)
3. computing KL(q ∥ p)\mathrm{KL}(q \,\Vert\, p) when you meant KL(p ∥ q)\mathrm{KL}(p \,\Vert\, q)the right number only on symmetric examplestest_kl_is_not_symmetric (mutant s03)
4. np.exp(r) - r - 1rounding noise instead of r2/2r^2/2 for small rrtest_kl_k3_small_r_precision (mutant s07)
5. -sum(p * log(p)) with p=softmax(z)p = \mathrm{softmax}(z)NaN entropy as soon as one probability underflowstest_entropy_from_logits (mutant s10)
6. np.log2 instead of the natural logevery number off by a factor of ln⁡2\ln 2test_hand_example (mutant s11)
7. the k3 ratio upside down, r=log⁡π−log⁡πrefr = \log \pi - \log \pi_{\mathrm{ref}}still nonnegative, but its mean is not KLtest_kl_k3_mean_is_kl (mutant s08)
8. letting −∞−(−∞)-\infty - (-\infty) or 0⋅∞0 \cdot \infty throughNaN where the answer is 0 or +∞+\inftytest_kl_from_logprobs_infinite_cases (mutants s06, s13)

| Forward | L12.2 | Registered module relationship. | | Forward | L12.4 | Registered module relationship. | | Forward | M11.4 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackM09.2log_softmax inside entropy_from_logits; the max shift behind every stable log-probability
ForwardM08.3its cross-entropy VJP is checked as the gradient of this cross_entropy
ForwardL0.3the training loss is H(onehot,softmax(z))H(\text{onehot}, \mathrm{softmax}(z)), fused and stable
ForwardL8.1logs entropy_from_logits per sampling step
ForwardL8.6speculative decoding accepts with probability 1−TV(p,q)1 - \mathrm{TV}(p, q); JS and KL frame the same comparison
ForwardL12.3the KL penalty to the reference model, per token, with kl_k3
ForwardM11.2perplexity is eHe^{H} of the per-token cross-entropy

If you skip this module, ol check M08.3 stops with M08.3 needs M11.1: build it, or rerun with --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
entropy, klscipy.stats.entropya base argument and normalization of its inputsscipy/stats/_entropy.py
kl_from_logprobstorch.nn.functional.kl_divtakes log qq as input and pp as target (note the order), log_target, batch reductionstorch/nn/functional.py
kl_k3TRL’s GRPO trainerthe per-token KL penalty er−r−1e^{r} - r - 1 against the reference modeltrl/trainer/grpo_trainer.py
jsscipy.spatial.distance.jensenshannonreturns the JS distance (the square root), with a basescipy/spatial/distance.py