応用数学

分布間ダイバージェンス

2つの確率分布の違いを測る量を、KLの向きとJSの性質から使い分けます。

  • B|標準
  • 応用数学

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

ひとことで言うと

ダイバージェンスは、正解側の分布 pp と近似側の分布 qq の違いを、確率の置き方に応じて数値化します。KLは向きを持つため、単なる「距離」として扱えません。分布の台がずれると、片方の向きは有限でも他方は発散します。

地図 pp と実際に歩くルート qq の比較で、全地点を覆うルートを選ぶのか、歩く先を絞って危険地点を避けるのかで評価基準が変わる、と考えると向きの違いを捉えやすいです。

なぜ必要か

平均値や個別サンプルの誤差だけでは、分布全体の形、特に複数の山(モード)を比較できません。近似分布を学習するときは「正解分布が確率を置く場所を漏らさない」のか、「正解分布が低確率な場所には置かない」のかを先に決めます。その選択がKLの向きに現れます。

比較で優先すること選びやすい向き
正解分布の山をできるだけ覆うpp から qq へのKL
低確率領域へのはみ出しを抑えるqq から pp へのKL

仕組み

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) が小さいと強く罰します。したがって DKL(p∥q)D_{\mathrm{KL}}(p\|q) は複数モードを覆う方向(モード平均化)になりやすい一方、

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は非負ですが、対称性 D(p∥q)=D(q∥p)D(p\|q)=D(q\|p) と三角不等式を満たさないため距離ではありません。さらに p(x)>0p(x)>0 なのに q(x)=0q(x)=0 となる点があれば、順方向の項は発散します。

JSダイバージェンスは混合分布 m=(p+q)/2m=(p+q)/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)

と定義します。pp と qq を入れ替えても同じで、対称かつ有界です(底が2の対数なら 00 から 11)。両方の台を含む mm を使うため、KLのような片側のゼロでの発散も避けやすくなります。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
KLの性質非負だが非対称、距離ではない対称性や三角不等式を足す
KLの向きp∥qp\|q はモード平均化、q∥pq\|p はモード探索の傾向向きを入れ替えて同じ結果とする
JSとの比較JSは対称・有界で、混合分布を使うJSも片側のKLだとする
台の扱いp>0,q=0p>0,q=0 の順方向KLは発散「確率0の項は常に無視」とする

実装で確かめる

NumPyで離散分布を計算します。ゼロ確率を含むときは、数学上の発散を無理に log(0) へ渡さないよう、寄与を条件分岐します。

import numpy as np

p = np.array([0.5, 0.5, 0.0])
q = np.array([0.5, 0.0, 0.5])

def kl(a, b):
    if np.any((a > 0) & (b == 0)):
        return np.inf
    mask = a > 0
    return np.sum(a[mask] * np.log(a[mask] / b[mask]))

m = (p + q) / 2
js = (kl(p, m) + kl(q, m)) / 2
print(kl(p, q), kl(q, p), js)

実行結果は inf inf 0.34657359027997264 です。JSは有限ですが、これはKLの向きを消したのではなく、共通の混合分布を挟んで両方向を平均した結果です。

取り違えやすいもの

量使い分け
DKL(p∥q)D_{\mathrm{KL}}(p\|q)基準分布の高確率領域を近似側に覆わせたいとき
DKL(q∥p)D_{\mathrm{KL}}(q\|p)近似側が基準分布の低確率領域へ出ないことを重視するとき
DJS(p,q)D_{\mathrm{JS}}(p,q)対称で有界な比較量が必要なとき
ℓ2ℓ_2 距離確率の比や対数ではなく、座標ごとの差を測るとき

GANでは、生成分布とデータ分布の比較にJSを含む目的関数との接続がありますが、ここで重要なのは「どのダイバージェンスを最小化しているか」を目的関数ごとに確認することです。なお、PyTorchの KLDivLoss は入力を対数確率で受け、数学上のバッチ平均には reduction="batchmean" を使います。

想起チェック

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

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

KLの向きを入れ替えると何が変わるか

前者は複数モードを平均化しやすく、後者は一つのモードを探索しやすいです。これはサンプルされる分布が違うためです。

JSダイバージェンスの性質は

対称かつ有界で、混合分布 mm を介した2つのKLの平均です。

出典