応用数学

KLとJSダイバージェンス

KLダイバージェンスの向きが生む違いと、対称化したJSダイバージェンスの性質を式と実装で整理します。

  • B|標準
  • 応用数学

数式の記号で止まったら 記号の読み方 (∂・⊙・転置・上付き添字を、読み方から)

ひとことで言うと

KLダイバージェンスは、基準分布 pp を近似分布 qq で表したときの情報量のずれです。JSダイバージェンスは、pp と qq の中間分布を介してKLを対称化した量です。どちらも「確率分布同士の差」を測りますが、KLは向きを持ち、JSは対称です。

KLは「正解分布が出す場所を、近似分布がどれだけ説明し損ねたか」を片方向に採点します。JSは両者の中間案を作ってから、正解側と近似側を同じ重みで採点します。

なぜ必要か

平均二乗誤差は、確率分布の形や確率の比を直接扱いません。たとえば、正解分布が確率を置く領域で q(x)q(x) がゼロなら、そこを説明できないことを大きく罰したい場合があります。このとき確率の対数比を使うKLが自然です。

一方、KLは pp と qq を入れ替えると値が変わります。分布間の比較を一方に依存しない形で行いたいときは、混合分布を使うJSを選びます。したがって「どちらが正しい距離か」ではなく、何を基準に誤差を測るかで使い分けます。

目的向いている量
基準分布の確率質量を近似側に覆わせるDKL(p∥q)D_{\mathrm{KL}}(p\|q)
比較を対称にし、上限のある値にするDJS(p,q)D_{\mathrm{JS}}(p,q)

仕組み

p(x)p(x) を基準分布、q(x)q(x) を近似分布、xx を確率変数とすると、KLダイバージェンスは

DKL(p∥q)=Ex∼p[log⁡p(x)q(x)]D_{\mathrm{KL}}(p\|q)=\mathbb{E}_{x\sim p}\left[\log\frac{p(x)}{q(x)}\right]

です。期待値を pp から取るため、p(x)p(x) が大きい場所で q(x)q(x) が小さいと強く効きます。pp と qq の順序を逆にした

DKL(q∥p)=Ex∼q[log⁡q(x)p(x)]D_{\mathrm{KL}}(q\|p)=\mathbb{E}_{x\sim q}\left[\log\frac{q(x)}{p(x)}\right]

は、qq が確率を置いた場所で pp を評価します。これが同じ分布対でも値や最適化の傾向が変わる理由です。KLは常に非負ですが、一般には対称でなく、三角不等式も満たさないため距離ではありません。p(x)>0p(x)>0 なのに q(x)=0q(x)=0 なら、DKL(p∥q)D_{\mathrm{KL}}(p\|q) は発散します。

JSでは混合分布 m(x)=(p(x)+q(x))/2m(x)=(p(x)+q(x))/2 を作り、

DJS(p,q)=12DKL(p∥m)+12DKL(q∥m)D_{\mathrm{JS}}(p,q)=\frac{1}{2}D_{\mathrm{KL}}(p\|m)+\frac{1}{2}D_{\mathrm{KL}}(q\|m)

とします。mm は両方の分布が確率を置く場所を含むので、片方の分布のゼロ確率による直接的な発散を避けられます。JSは pp と qq を交換しても同じです。対数の底を2にすれば値は 00 以上 11 以下になり、自然対数なら上限は log⁡2\log 2 です。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
KLの性質非負だが、非対称で距離ではない非負だから距離だとする
KLの向き期待値をどちらの分布から取るかで変わるDKL(p∥q)=DKL(q∥p)D_{\mathrm{KL}}(p\|q)=D_{\mathrm{KL}}(q\|p) とする
JSの定義m=(p+q)/2m=(p+q)/2 を介した2つのKLの平均KLを単に絶対値化した量とする
上限対数の底が2なら 11、自然対数なら log⁡2\log 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はゼロ確率で発散しますが、混合分布 mm とのKLは有限になるため、JSは有限です。

取り違えやすいもの

量切り分け
DKL(p∥q)D_{\mathrm{KL}}(p\|q)pp を基準に qq の説明不足を測る。順序を入れ替えられない
DKL(q∥p)D_{\mathrm{KL}}(q\|p)qq を基準に評価する別の量。同じ「KL」でも値は別
DJS(p,q)D_{\mathrm{JS}}(p,q)混合分布を介する対称な比較。上限も対数の底で決まる
交差エントロピーH(p,q)=−Ep[log⁡q(x)]H(p,q)=-\mathbb{E}_{p}[\log q(x)]。H(p,q)=H(p)+DKL(p∥q)H(p,q)=H(p)+D_{\mathrm{KL}}(p\|q) で、KLそのものではない

特に実装では、損失の入力が確率なのか対数確率なのかを確認します。数式の qq をそのまま渡すAPIとは限らず、対数確率を受け取る実装では、式の形を変換してから計算します。

想起チェック

KLダイバージェンスは距離か

距離ではありません。非負ですが、一般に対称性と三角不等式を満たしません。

KLで順序が重要なのはなぜか

期待値を取る分布が変わるからです。DKL(p∥q)D_{\mathrm{KL}}(p\|q) は pp の確率質量を基準に評価します。

JSで使う混合分布は何か

m=(p+q)/2m=(p+q)/2 です。pp と qq それぞれとのKLを平均することで、対称な量にします。

出典