Entropy, cross-entropy, KL, JS, and the k3 estimator
Overview
Section titled “Overview”| Module | M11.1 · build · Python · Pass 2 · 2 to 3 h |
| You build | python/tinyllm/info/entropy.py: entropy, cross_entropy, kl, kl_from_logprobs, js, kl_k3, entropy_from_logits |
| Contract | course/contracts/py/tinyllm/info/entropy.pyi |
| Tests | course/tests/M11.1/ (what they check: section 4) |
| Needs | M09.2 stable numerics (log_softmax, or --ref-deps) |
| Used by | M08.3 differentiates this cross-entropy · L0.3 the training loss · later: L8.1 per-step entropy logging, L8.6 acceptance rate as , L12.3 the KL penalty with kl_k3, M11.2 perplexity · later: L12.2, L12.4, M11.4 |
| Milestone | MS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains) |
| Optional depth | Cover 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) |
Key Takeaways
Section titled “Key Takeaways”- Entropy is the average surprise of a distribution, between 0 (certain) and (uniform) (
test_entropy_bounds). - Cross-entropy splits as : the training loss is the data’s own entropy, which no model can beat, plus the model’s excess (
test_cross_entropy_decomposes). - with equality only at (Gibbs), and it is not symmetric; Jensen-Shannon is symmetric and bounded by (
test_gibbs_kl_nonnegative,test_kl_is_not_symmetric,test_js_properties). - From model outputs, compute KL with log-probabilities, never with ; the k3 estimator estimates KL from samples without bias and is never negative (
test_kl_from_logprobs_extreme_logits,test_kl_k3_mean_is_kl).
How to work this chapter
Section titled “How to work this chapter”ol start M11.1 # stubs entropy.py into your repo, contract alongsideol tests M11.1 # read the test catalog first: rung R0, you write no tests hereol check M11.1 # exit code is the verdictol check M11.1 --ref-deps # only if your M09.2 is not passing yetol diff M11.1 # after passing: your code against the reference1. Why now
Section titled “1. Why now”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.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| number of outcomes (vocabulary size for a language model) | int | |
| probability distributions over the outcomes: , | float64[n] | |
| the natural logarithm; results are in nats (divide by for bits) | ||
| entropy of | scalar | |
| cross-entropy of relative to | scalar | |
| Kullback-Leibler divergence from to | scalar, | |
| the mixture of and | float64[n] | |
| Jensen-Shannon divergence | scalar in | |
a policy and a reference model (L12.3) | distributions | |
| an outcome sampled from | index | |
| log ratio at a sample | scalar | |
| the k3 estimator of | scalar, | |
logits, with a distribution (M09.2) | float64[n] |
Surprise. An outcome of probability carries surprise : a certain outcome () carries none, a rare one a lot, and the surprise of two independent outcomes adds, because . With the unit is the bit: is the length of the best code word for outcome .
Entropy is the average surprise.
A term with is defined as 0, because as : an outcome that never happens costs nothing. since every surprise is , and , 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 and you predict with , your average surprise is
This is the language-model loss. With the one-hot distribution of the observed next token , it is , and averaging over positions gives the negative log-likelihood your bigram already prints. A term with and is : the model said impossible, and it happened.
KL is the excess, and it is never negative. Subtract the entropy:
Proof that it is (Gibbs’ inequality): for , with equality only at . Over the outcomes with ,
Equality needs everywhere. So minimizing the cross-entropy over minimizes KL, and the loss floor is : the entropy of the data itself. Training a language model is pushing toward 0.
KL has a direction. weighs the log ratio by , so it punishes for being small where is large (it is “mass covering”), and is infinite if misses any outcome of . punishes the opposite and prefers a that sits on one mode of . They are different numbers (section 4 has one example), and L12 picks one on purpose.
KL from log-probabilities. Models output log-probabilities (), which are finite even where the probability underflows. Write
a difference of moderate numbers, and never form , which is or when the logits are large. Two edge cases: means and contributes 0 (even against , where the difference would be NaN); a finite against is , even when underflows to 0 (where would be NaN).
Jensen-Shannon is a symmetric, bounded cousin. Compare both with their average :
is positive wherever or is, so JS is always finite, and since , every : , reached by distributions with disjoint supports.
Estimating KL from samples: k3. When is a vocabulary and you only have the log-probabilities of the tokens you sampled, you estimate by averaging over samples. The plain estimator is unbiased but negative for many samples. Schulman’s k3 adds a term with mean zero:
so has the same mean, (the sum runs over ‘s support, which is everything for softmax outputs). And every sample is , because the exponential lies above its tangent at 0: .
k3 for small . Near , is tiny while . Computing subtracts numbers near 1 and keeps only the rounding error of , about , against a true value of at . np.expm1(r) computes directly to full relative precision, so keeps the answer. This is exactly the regime at the start of post-training, when the policy has barely moved.
Entropy from logits. With from M09.2, , with masked entries () contributing 0. Computing and then meets as soon as one probability underflows.
3. Worked example by hand
Section titled “3. Worked example by hand”Take , , . Every log is a multiple of except one.
Entropy. nats, which is 1.5 bits.
Cross-entropy. .
KL. Directly: . Check: . Same number.
Jensen-Shannon. .
| 0 | ||
So ; by symmetry of this example is the same, and , well under .
k3. Let and . The log ratios are :
| 1 | 1/2 | ||
| 2 | 1/4 | 0 | 0 |
| 3 | 1/4 |
The mean under : . Exactly , as promised, and every row is nonnegative.
Entropy from logits. gives , softmax , so the entropy is again .
These numbers are the first case in section 4, test_hand_example.
4. The interface
Section titled “4. The interface”# python/tinyllm/info/entropy.py (natural logs: nats)def entropy(p: ArrayLike, axis: int = -1) -> NDArraydef cross_entropy(p: ArrayLike, q: ArrayLike, axis: int = -1) -> NDArraydef kl(p: ArrayLike, q: ArrayLike, axis: int = -1) -> NDArraydef kl_from_logprobs(logp: ArrayLike, logq: ArrayLike, axis: int = -1) -> NDArraydef js(p: ArrayLike, q: ArrayLike) -> NDArray # last axisdef kl_k3(logp_ref: ArrayLike, logp: ArrayLike) -> NDArray # elementwise, no reductiondef entropy_from_logits(z: ArrayLike, axis: int = -1) -> NDArrayEach 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 . Compute the zero conventions with np.where on the terms; adding a small epsilon inside the logs biases every answer and turns a true into a large finite number.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | every section 3 number, in nats | you and the test agree on the definitions |
test_matches_scipy_golden | golden | scipy’s entropy, KL, and JS on dense, sparse, peaked, and column-wise inputs | an independent implementation agrees |
test_gibbs_kl_nonnegative | property | on 200 random pairs, exactly 0 at | a KL penalty never rewards drift (L12.3) |
test_cross_entropy_decomposes | property | reading the loss: floor plus excess | |
test_kl_is_not_symmetric | unit | vs uniform: 0.36806 one way, 0.51083 the other | forward and reverse KL are different objectives |
test_zero_probability_conventions | boundary | , , no NaN | sparse distributions, masked tokens |
test_rejects_negative_probabilities | boundary | a negative entry is a ValueError everywhere | a logit passed as a probability fails loudly |
test_entropy_bounds | property | 0 for one-hot, for uniform, in between otherwise | reading L8.1’s entropy log |
test_axis_and_batch_shapes | unit | [B, V] reduces to [B]; axis=0 reduces columns | batched losses |
test_js_properties | property | symmetric, in , for disjoint supports | a bounded comparison of two models |
test_kl_from_logprobs_matches_kl | differential | equals kl on the same distributions, entries included | the form models use |
test_kl_from_logprobs_extreme_logits | boundary | log-probabilities from logits near give a finite, correct KL | late training, low temperature |
test_kl_from_logprobs_infinite_cases | boundary | against is 0; finite (even ) against is | no NaN from or |
test_kl_k3_is_nonnegative | property | on 1000 samples, 0 at , at | per-sample penalties are never negative |
test_kl_k3_small_r_precision | boundary | relative error below at | the start of L12.3, when the policy barely moved |
test_kl_k3_mean_is_kl | statistical | exact enumeration equals KL; 20000 samples within 4 standard errors | the estimator is unbiased |
test_entropy_from_logits | boundary | equals , finite at , masks count as 0 | L8.1 logs it from logits |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
1. p * np.log(p) with , or an epsilon inside the logs | NaN entropy for sparse rows; a finite KL where it is infinite | test_zero_probability_conventions (mutants s01, s02) |
| 2. forming from exponentiated log-probabilities | 0/0 and inf/inf once logits are large | test_kl_from_logprobs_extreme_logits (mutant s05) |
| 3. computing when you meant | the right number only on symmetric examples | test_kl_is_not_symmetric (mutant s03) |
4. np.exp(r) - r - 1 | rounding noise instead of for small | test_kl_k3_small_r_precision (mutant s07) |
5. -sum(p * log(p)) with | NaN entropy as soon as one probability underflows | test_entropy_from_logits (mutant s10) |
6. np.log2 instead of the natural log | every number off by a factor of | test_hand_example (mutant s11) |
| 7. the k3 ratio upside down, | still nonnegative, but its mean is not KL | test_kl_k3_mean_is_kl (mutant s08) |
| 8. letting or through | NaN where the answer is 0 or | test_kl_from_logprobs_infinite_cases (mutants s06, s13) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Forward | L12.2 | Registered module relationship. |
| Forward | L12.4 | Registered module relationship. |
| Forward | M11.4 | Registered call site uses this module. |
| Direction | Module | How it uses this |
|---|---|---|
| Back | M09.2 | log_softmax inside entropy_from_logits; the max shift behind every stable log-probability |
| Forward | M08.3 | its cross-entropy VJP is checked as the gradient of this cross_entropy |
| Forward | L0.3 | the training loss is , fused and stable |
| Forward | L8.1 | logs entropy_from_logits per sampling step |
| Forward | L8.6 | speculative decoding accepts with probability ; JS and KL frame the same comparison |
| Forward | L12.3 | the KL penalty to the reference model, per token, with kl_k3 |
| Forward | M11.2 | perplexity is 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.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
entropy, kl | scipy.stats.entropy | a base argument and normalization of its inputs | scipy/stats/_entropy.py |
kl_from_logprobs | torch.nn.functional.kl_div | takes log as input and as target (note the order), log_target, batch reductions | torch/nn/functional.py |
kl_k3 | TRL’s GRPO trainer | the per-token KL penalty against the reference model | trl/trainer/grpo_trainer.py |
js | scipy.spatial.distance.jensenshannon | returns the JS distance (the square root), with a base | scipy/spatial/distance.py |