KL Divergence
Two models divide probabilities differently among the same answers. One assigns 90% to A and 10% to B; the other gives each 50%. You want to measure how well the second reproduces the first model's predictions.
KL Divergence compares two probability distributions, sets of probabilities for the same events. You must specify the reference distribution and the one approximating it. Each event's contribution depends on its reference probability and the ratio of the two probabilities.
For these numbers, comparing the first to the second gives about 0.368, and the reverse about 0.511, using natural logarithms. Direction matters: this is not a symmetric distance like the number of kilometers between cities.
In Knowledge Distillation, teacher and student predictions can be compared. A small KL indicates agreement, rather than the teacher's truthfulness. Assigning zero to an event with positive reference probability gives a mathematically infinite penalty.
Mechanism and details
and are probabilities of the same event , such as a particular token. The formula and its relationship with Cross-entropy are given in the SciPy: entropy documentation. We have : the entropy of the reference distribution is subtracted from Cross-entropy. When is fixed, this difference does not change the minimum with respect to the parameters of the model producing .
The direction of comparison matters
An original example: for and , the result is about 0.368. After swapping the distributions' roles, the result is about 0.511. We use the natural logarithm, so the unit is nats. KL Divergence is not a symmetric distance: the distribution on each side must always be specified.
If but , the result is infinite. A term with is taken to be zero. These edge cases are defined by SciPy: rel_entr.
In Knowledge Distillation, the teacher's prediction can serve as the reference distribution and the student's as the approximation. A small value indicates agreement between these distributions on the evaluated data, not the truth of the teacher's response.
When implementing this, care is needed with argument order. PyTorch: KLDivLoss expects model log-probabilities as input and the reference distribution as target when log_target=False. For an “examples × classes” matrix, batchmean reduction sums terms over classes and averages over examples; mean averages all elements, producing a different scale.