ひとことで言うと
KLダイバージェンスは、基準分布 を近似分布 で表したときの情報量のずれです。JSダイバージェンスは、 と の中間分布を介してKLを対称化した量です。どちらも「確率分布同士の差」を測りますが、KLは向きを持ち、JSは対称です。
KLは「正解分布が出す場所を、近似分布がどれだけ説明し損ねたか」を片方向に採点します。JSは両者の中間案を作ってから、正解側と近似側を同じ重みで採点します。
なぜ必要か
平均二乗誤差は、確率分布の形や確率の比を直接扱いません。たとえば、正解分布が確率を置く領域で がゼロなら、そこを説明できないことを大きく罰したい場合があります。このとき確率の対数比を使うKLが自然です。
一方、KLは と を入れ替えると値が変わります。分布間の比較を一方に依存しない形で行いたいときは、混合分布を使うJSを選びます。したがって「どちらが正しい距離か」ではなく、何を基準に誤差を測るかで使い分けます。
| 目的 | 向いている量 |
|---|---|
| 基準分布の確率質量を近似側に覆わせる | |
| 比較を対称にし、上限のある値にする |
仕組み
を基準分布、 を近似分布、 を確率変数とすると、KLダイバージェンスは
です。期待値を から取るため、 が大きい場所で が小さいと強く効きます。 と の順序を逆にした
は、 が確率を置いた場所で を評価します。これが同じ分布対でも値や最適化の傾向が変わる理由です。KLは常に非負ですが、一般には対称でなく、三角不等式も満たさないため距離ではありません。 なのに なら、 は発散します。
JSでは混合分布 を作り、
とします。 は両方の分布が確率を置く場所を含むので、片方の分布のゼロ確率による直接的な発散を避けられます。JSは と を交換しても同じです。対数の底を2にすれば値は 以上 以下になり、自然対数なら上限は です。
試験でどう問われるか
| 問われ方 | 正解に寄る条件 | 引っかけ |
|---|---|---|
| KLの性質 | 非負だが、非対称で距離ではない | 非負だから距離だとする |
| KLの向き | 期待値をどちらの分布から取るかで変わる | とする |
| JSの定義 | を介した2つのKLの平均 | KLを単に絶対値化した量とする |
| 上限 | 対数の底が2なら 、自然対数なら | 底を無視して上限を固定する |
実装で確かめる
離散分布で、KLの向きとJSを計算します。ゼロ確率の項は数学上の発散を log(0) に任せず、条件分岐で表します。
import math
p = [0.5, 0.5, 0.0]
q = [0.5, 0.0, 0.5]
def kl(a, b):
if any(x > 0 and y == 0 for x, y in zip(a, b)):
return math.inf
return sum(x * math.log(x / y) for x, y in zip(a, b) if x > 0)
m = [(x + y) / 2 for x, y in zip(p, q)]
js = (kl(p, m) + kl(q, m)) / 2
print(kl(p, q), kl(q, p), js)
出力は inf inf 0.34657359027997264 です。両方向のKLはゼロ確率で発散しますが、混合分布 とのKLは有限になるため、JSは有限です。
取り違えやすいもの
| 量 | 切り分け |
|---|---|
| を基準に の説明不足を測る。順序を入れ替えられない | |
| を基準に評価する別の量。同じ「KL」でも値は別 | |
| 混合分布を介する対称な比較。上限も対数の底で決まる | |
| 交差エントロピー | 。 で、KLそのものではない |
特に実装では、損失の入力が確率なのか対数確率なのかを確認します。数式の をそのまま渡すAPIとは限らず、対数確率を受け取る実装では、式の形を変換してから計算します。
想起チェック
KLダイバージェンスは距離か
距離ではありません。非負ですが、一般に対称性と三角不等式を満たしません。
KLで順序が重要なのはなぜか
期待値を取る分布が変わるからです。 は の確率質量を基準に評価します。
JSで使う混合分布は何か
です。 と それぞれとのKLを平均することで、対称な量にします。