この 章の 目標
エントロピー・交差エントロピー・KL ダイバージェンス・相互情報量の 基本性質と ギブスの 不等式を 証明し、最尤推定が 経験分布との KL ダイバージェンスの 最小化である ことを 示せる
混合ガウスモデルの EM アルゴリズムを 導いて 対数尤度の 単調性を 証明し、局所解と 尤度の 非有界性と いう 限界を 説明できる
log p ( x ) = ELBO ( q ) + KL ( q ∥ p ( ⋅ ∣ x ) ) \log p(x) = \operatorname{ELBO}(q) + \operatorname{KL}(q \Vert p(\cdot \mid x)) log p ( x ) = ELBO ( q ) + KL ( q ∥ p ( ⋅ ∣ x )) を 導き、平均場近似の 更新式と、EM が 変分推論の 特別な 場合である ことを 説明できる
隠れマルコフモデルの 前向きアルゴリズムを 導き、計算量が O ( T K 2 ) O(TK^2) O ( T K 2 ) である ことを 示せる
変分オートエンコーダ・拡散モデル・ 大規模言語モデルを、確率モデルと 学習の 目的関数の 言葉で 読める
前提 :第1章 、22-statistics 第1章 (条件付き分布・ 多変量正規分布)、 22-statistics 第3章 (最尤推定)。7.5 節では 11-probability 第6章 の マルコフ連鎖を、7.6 節では 第2章の ソフトマックス回帰と 第5章の ニューラルネットワークを 使う。 22-statistics 第7章 (事後分布・マルコフ連鎖モンテカルロ法)と 比べると よい。
工場の センサーの 値が ふだんの 分布から 外れたら 異常を 疑う。購入額の 分布に 山が 2 つ あれば、性質の 違う 客が 混ざっているのだろう。設備の 劣化は 直接 見えず、アラームの 記録から 推し量るしかない。いずれも、データ x x x の 分布 p ( x ) p(x) p ( x ) を、観測できない 潜在変数 (latent variable) も 含めて モデル化する 問題である。データの 生まれ方を 確率分布で 記述し、新しい データを 生成できる モデルを 生成モデル (generative model) と いう。
本章では KL ダイバージェンスと 最尤推定(7.1・7.2 節)、EM アルゴリズム(7.3 節)、変分推論(7.4 節)、隠れマルコフモデル(7.5 節)を 学び、近年の 生成モデルを 同じ 言葉で 読む(7.6 節)。 log \log log は 自然対数、「分布 p p p 」は 確率関数または 密度を 指し、和は 連続の 場合は 積分に 読み替える。測度論的な 細部は 11-probability に 譲る。
7.1 エントロピーと KL ダイバージェンス
定義 7.1 (エントロピー・交差エントロピー・KL ダイバージェンス)p , q p, q p , q を 有限集合または 可算集合上の 確率分布とし、 A = { x ∣ p ( x ) > 0 } A = \lbrace x \mid p(x) > 0 \rbrace A = { x ∣ p ( x ) > 0 } と する。
H ( p ) = − ∑ x ∈ A p ( x ) log p ( x ) , H ( p , q ) = − ∑ x ∈ A p ( x ) log q ( x ) , KL ( p ∥ q ) = ∑ x ∈ A p ( x ) log p ( x ) q ( x ) H(p) = -\sum_{x \in A} p(x)\log p(x), \qquad H(p, q) = -\sum_{x \in A} p(x)\log q(x), \qquad \operatorname{KL}(p \Vert q) = \sum_{x \in A} p(x)\log\frac{p(x)}{q(x)} H ( p ) = − x ∈ A ∑ p ( x ) log p ( x ) , H ( p , q ) = − x ∈ A ∑ p ( x ) log q ( x ) , KL ( p ∥ q ) = x ∈ A ∑ p ( x ) log q ( x ) p ( x )
を それぞれ エントロピー (entropy)、交差エントロピー (cross-entropy)、KL ダイバージェンス (カルバック–ライブラー情報量, Kullback–Leibler divergence)と いう。 q ( x ) = 0 q(x) = 0 q ( x ) = 0 と なる x ∈ A x \in A x ∈ A が あれば後の 2 つは ∞ \infty ∞ と する。密度でも 和を A A A 上の 積分に 変えて 同様に 定める( A A A 上で q = 0 q = 0 q = 0 と なる 部分の 確率が 正なら後の 2 つは ∞ \infty ∞ )。X ∼ p X \sim p X ∼ p の とき H ( X ) = H ( p ) H(X) = H(p) H ( X ) = H ( p ) とも 書く。
− log p ( x ) -\log p(x) − log p ( x ) は 起こりにくい値ほど 大きい「驚き」で、エントロピーは その 平均(分布の 不確かさ)である。KL ダイバージェンスの 項は 負にも なるが、下の log u ≤ u − 1 \log u \leq u - 1 log u ≤ u − 1 より p log p q ≥ p − q ≥ − q p\log\frac{p}{q} \geq p - q \geq -q p log q p ≥ p − q ≥ − q なので、負の 項の 絶対値の 和は 1 1 1 以下であり、和は ( − ∞ , ∞ ] (-\infty, \infty] ( − ∞ , ∞ ] の 値と して 定まる。 H ( p ) < ∞ H(p) < \infty H ( p ) < ∞ なら H ( p , q ) = H ( p ) + KL ( p ∥ q ) H(p, q) = H(p) + \operatorname{KL}(p \Vert q) H ( p , q ) = H ( p ) + KL ( p ∥ q ) である。密度の エントロピー(微分エントロピー)は 負にもなる。
定理 7.2 (ギブスの 不等式, Gibbs' inequality) KL ( p ∥ q ) ≥ 0 \operatorname{KL}(p \Vert q) \geq 0 KL ( p ∥ q ) ≥ 0 であり、等号は p = q p = q p = q の ときに 限る。密度の 場合も KL ( p ∥ q ) ≥ 0 \operatorname{KL}(p \Vert q) \geq 0 KL ( p ∥ q ) ≥ 0 で、等号は p p p と q q q が ほとんど 至る ところ(ルベーグ測度 0 0 0 の 集合を 除いて)等しい ときに 限る。
証明. g ( u ) = u − 1 − log u g(u) = u - 1 - \log u g ( u ) = u − 1 − log u (u > 0 u > 0 u > 0 )は g ′ ( u ) = 1 − 1 / u g'(u) = 1 - 1/u g ′ ( u ) = 1 − 1/ u が u = 1 u = 1 u = 1 の 前後で 負から 正に 変わるので、最小値 g ( 1 ) = 0 g(1) = 0 g ( 1 ) = 0 を とる。すな わち log u ≤ u − 1 \log u \leq u - 1 log u ≤ u − 1 で、等号は u = 1 u = 1 u = 1 に 限る。 q = 0 q = 0 q = 0 と なる A A A の 点が あれば KL = ∞ \operatorname{KL} = \infty KL = ∞ なので、A A A 上 q > 0 q > 0 q > 0 と する。 p log q p = ( q − p ) − p ⋅ g ( q / p ) p\log\frac{q}{p} = (q - p) - p \cdot g(q/p) p log p q = ( q − p ) − p ⋅ g ( q / p ) より
− KL ( p ∥ q ) = ∑ x ∈ A ( q ( x ) − p ( x ) ) − ∑ x ∈ A p ( x ) g ( q ( x ) p ( x ) ) ≤ ∑ x ∈ A q ( x ) − 1 ≤ 0 -\operatorname{KL}(p \Vert q) = \sum_{x \in A}\bigl(q(x) - p(x)\bigr) - \sum_{x \in A}p(x)\,g\Bigl(\frac{q(x)}{p(x)}\Bigr) \leq \sum_{x \in A}q(x) - 1 \leq 0 − KL ( p ∥ q ) = x ∈ A ∑ ( q ( x ) − p ( x ) ) − x ∈ A ∑ p ( x ) g ( p ( x ) q ( x ) ) ≤ x ∈ A ∑ q ( x ) − 1 ≤ 0
である(第 1 の 和は 絶対収束し、第 2 の 和の 項は 0 0 0 以上)。等号なら 第 2 の 和は 0 0 0 で ∑ x ∈ A q ( x ) = 1 \sum_{x \in A}q(x) = 1 ∑ x ∈ A q ( x ) = 1 だから、A A A 上で q = p q = p q = p 、A A A の 外で q = 0 = p q = 0 = p q = 0 = p である。密度の 場合も、 A A A 上で q = 0 q = 0 q = 0 と なる 部分は(確率が 正なら KL = ∞ \operatorname{KL} = \infty KL = ∞ なので)零集合と して 除けば 同じ 式が 成り立ち、等号なら 0 0 0 以上の 関数 p ⋅ g ( q / p ) p \cdot g(q/p) p ⋅ g ( q / p ) の A A A 上の 積分が 0 0 0 なので A A A 上ほとんど 至る ところ q = p q = p q = p で(06-measure-integration 第3章 系 3.12 の 4)、∫ A q = 1 \int_A q = 1 ∫ A q = 1 から A A A の 外でも ほとんど 至る ところ q = 0 q = 0 q = 0 である。□ \square □
系 7.3
H ( p ) < ∞ H(p) < \infty H ( p ) < ∞ なら H ( p , q ) ≥ H ( p ) H(p, q) \geq H(p) H ( p , q ) ≥ H ( p ) で、等号は q = p q = p q = p に 限る。すな わち q ↦ H ( p , q ) q \mapsto H(p, q) q ↦ H ( p , q ) は q = p q = p q = p でだけ 最小に なる。
K K K 個の 元からなる 集合の 上では 0 ≤ H ( p ) ≤ log K 0 \leq H(p) \leq \log K 0 ≤ H ( p ) ≤ log K で、右の 等号は 一様分布に 限る。
証明. 1 は H ( p , q ) = H ( p ) + KL ( p ∥ q ) H(p, q) = H(p) + \operatorname{KL}(p \Vert q) H ( p , q ) = H ( p ) + KL ( p ∥ q ) と 定理 7.2 に よる。2 の 左は 各項が 0 0 0 以上であることにより、右は 一様分布 u u u に ついて KL ( p ∥ u ) = log K − H ( p ) \operatorname{KL}(p \Vert u) = \log K - H(p) KL ( p ∥ u ) = log K − H ( p ) と なることに よる。 □ \square □
例 7.4 (KL ダイバージェンスの 計算)表の 確率が 0.5 0.5 0.5 , 0.9 0.9 0.9 の ベルヌーイ分布を p , q p, q p , q と すると KL ( p ∥ q ) = 0.5 log 0.5 0.9 + 0.5 log 0.5 0.1 = log 5 3 ≈ 0.5108 \operatorname{KL}(p \Vert q) = 0.5\log\frac{0.5}{0.9} + 0.5\log\frac{0.5}{0.1} = \log\frac{5}{3} \approx 0.5108 KL ( p ∥ q ) = 0.5 log 0.9 0.5 + 0.5 log 0.1 0.5 = log 3 5 ≈ 0.5108 、KL ( q ∥ p ) = 0.9 log 1.8 + 0.1 log 0.2 ≈ 0.3681 \operatorname{KL}(q \Vert p) = 0.9\log 1.8 + 0.1\log 0.2 \approx 0.3681 KL ( q ∥ p ) = 0.9 log 1.8 + 0.1 log 0.2 ≈ 0.3681 で、向きに よって 値が 違う。正規分布どうしでは、密度の 比の 対数 log σ 2 σ 1 − ( x − μ 1 ) 2 2 σ 1 2 + ( x − μ 2 ) 2 2 σ 2 2 \log\frac{\sigma_2}{\sigma_1} - \frac{(x - \mu_1)^2}{2\sigma_1^2} + \frac{(x - \mu_2)^2}{2\sigma_2^2} log σ 1 σ 2 − 2 σ 1 2 ( x − μ 1 ) 2 + 2 σ 2 2 ( x − μ 2 ) 2 の 期待値を とり、 E [ ( X − μ 2 ) 2 ] = σ 1 2 + ( μ 1 − μ 2 ) 2 E[(X - \mu_2)^2] = \sigma_1^2 + (\mu_1 - \mu_2)^2 E [( X − μ 2 ) 2 ] = σ 1 2 + ( μ 1 − μ 2 ) 2 を 使って
KL ( N ( μ 1 , σ 1 2 ) ∥ N ( μ 2 , σ 2 2 ) ) = log σ 2 σ 1 + σ 1 2 + ( μ 1 − μ 2 ) 2 2 σ 2 2 − 1 2 (1) \operatorname{KL}\bigl(N(\mu_1, \sigma_1^2) \Vert N(\mu_2, \sigma_2^2)\bigr) = \log\frac{\sigma_2}{\sigma_1} + \frac{\sigma_1^2 + (\mu_1 - \mu_2)^2}{2\sigma_2^2} - \frac{1}{2} \tag{1} KL ( N ( μ 1 , σ 1 2 ) ∥ N ( μ 2 , σ 2 2 ) ) = log σ 1 σ 2 + 2 σ 2 2 σ 1 2 + ( μ 1 − μ 2 ) 2 − 2 1 ( 1 )
を 得る。例えば KL ( N ( 0 , 1 ) ∥ N ( 1 , 4 ) ) = log 2 − 1 4 ≈ 0.4431 \operatorname{KL}(N(0, 1) \Vert N(1, 4)) = \log 2 - \frac{1}{4} \approx 0.4431 KL ( N ( 0 , 1 ) ∥ N ( 1 , 4 )) = log 2 − 4 1 ≈ 0.4431 、逆向きは 2 − log 2 ≈ 1.3069 2 - \log 2 \approx 1.3069 2 − log 2 ≈ 1.3069 である(数値積分でも 確かめた)。共分散行列が 対角の 多変量正規分布どうしなら、成分ごとの (1) の 和に なる。KL ダイバージェンスは 対称でなく、三角不等式も 満たさない(問題 7.1)ので、距離ではない。
定義 7.5 (相互情報量)離散確率変数 X , Y X, Y X , Y の 同時分布を p ( x , y ) p(x, y) p ( x , y ) とし、X , Y X, Y X , Y を 独立に した 同時分布 p ( x ) p ( y ) p(x)p(y) p ( x ) p ( y ) を p X ⊗ p Y p_X \otimes p_Y p X ⊗ p Y と 書く。 I ( X ; Y ) = KL ( p ∥ p X ⊗ p Y ) I(X; Y) = \operatorname{KL}(p \Vert p_X \otimes p_Y) I ( X ; Y ) = KL ( p ∥ p X ⊗ p Y ) を 相互情報量 (mutual information)、H ( X ∣ Y ) = − ∑ x , y p ( x , y ) log p ( x ∣ y ) H(X \mid Y) = -\sum_{x, y}p(x, y)\log p(x \mid y) H ( X ∣ Y ) = − ∑ x , y p ( x , y ) log p ( x ∣ y ) を 条件付きエントロピーと いう(和は p ( x , y ) > 0 p(x, y) > 0 p ( x , y ) > 0 の 組に ついてとる)。
命題 7.6
I ( X ; Y ) = I ( Y ; X ) ≥ 0 I(X; Y) = I(Y; X) \geq 0 I ( X ; Y ) = I ( Y ; X ) ≥ 0 で、等号は X X X と Y Y Y が 独立の ときに 限る。
H ( X ) < ∞ H(X) < \infty H ( X ) < ∞ なら I ( X ; Y ) = H ( X ) − H ( X ∣ Y ) I(X; Y) = H(X) - H(X \mid Y) I ( X ; Y ) = H ( X ) − H ( X ∣ Y ) で、特に H ( X ∣ Y ) ≤ H ( X ) H(X \mid Y) \leq H(X) H ( X ∣ Y ) ≤ H ( X ) 。
証明. 1 は 定理 7.2 と、独立性が すべての x , y x, y x , y で p ( x , y ) = p ( x ) p ( y ) p(x, y) = p(x)p(y) p ( x , y ) = p ( x ) p ( y ) と なる ことと 同値である こと( 22-statistics 第1章 命題 1.4)に よる。2 は log p ( x , y ) p ( x ) p ( y ) = log p ( x ∣ y ) − log p ( x ) \log\frac{p(x, y)}{p(x)p(y)} = \log p(x \mid y) - \log p(x) log p ( x ) p ( y ) p ( x , y ) = log p ( x ∣ y ) − log p ( x ) の 期待値を とればよい( H ( X ) H(X) H ( X ) が 有限なので 和を 分けて よい)。 □ \square □
相互情報量は Y Y Y を 知る ことで 減る X X X の 不確かさである(問題 7.2)。相関係数と 違って、 0 0 0 に なるのは 独立な ときに 限るので、非線形な 依存も 捉える。例えば X X X が { − 1 , 0 , 1 } \lbrace -1, 0, 1 \rbrace { − 1 , 0 , 1 } 上の 一様分布で Y = X 2 Y = X^2 Y = X 2 なら、相関係数は 0 0 0 だが、Y Y Y は X X X で 決まるので H ( Y ∣ X ) = 0 H(Y \mid X) = 0 H ( Y ∣ X ) = 0 で、命題 7.6 を X , Y X, Y X , Y を 入れ替えて 使うと I ( X ; Y ) = H ( Y ) = log 3 − 2 3 log 2 ≈ 0.637 > 0 I(X; Y) = H(Y) = \log 3 - \frac{2}{3}\log 2 \approx 0.637 > 0 I ( X ; Y ) = H ( Y ) = log 3 − 3 2 log 2 ≈ 0.637 > 0 である。
7.2 最尤推定と KL ダイバージェンス
定理 7.7 (最尤推定と KL ダイバージェンス)データ x 1 , … , x n x_1, \dots, x_n x 1 , … , x n を パラメータ θ \theta θ の モデル p θ p_\theta p θ (22-statistics 第3章 の f ( x ; θ ) f(x; \theta) f ( x ; θ ) )で 表し、 ℓ ( θ ) = ∑ i log p θ ( x i ) \ell(\theta) = \sum_i\log p_\theta(x_i) ℓ ( θ ) = ∑ i log p θ ( x i ) と する。
データが 有限集合または 可算集合に 値を とる とき、 経験分布 p ^ n ( x ) = 1 n ∣ { i ∣ x i = x } ∣ \hat{p}_n(x) = \frac{1}{n}\lvert \lbrace i \mid x_i = x \rbrace \rvert p ^ n ( x ) = n 1 ∣{ i ∣ x i = x }∣ に ついて 1 n ℓ ( θ ) = − H ( p ^ n , p θ ) = − H ( p ^ n ) − KL ( p ^ n ∥ p θ ) \frac{1}{n}\ell(\theta) = -H(\hat{p}_n, p_\theta) = -H(\hat{p}_n) - \operatorname{KL}(\hat{p}_n \Vert p_\theta) n 1 ℓ ( θ ) = − H ( p ^ n , p θ ) = − H ( p ^ n ) − KL ( p ^ n ∥ p θ ) である。したがって、θ \theta θ が 対数尤度を 最大に する ことと KL ( p ^ n ∥ p θ ) \operatorname{KL}(\hat{p}_n \Vert p_\theta) KL ( p ^ n ∥ p θ ) を 最小に する ことは 同値である。
X ∼ p ∗ X \sim p^{\ast} X ∼ p ∗ で H ( p ∗ ) H(p^{\ast}) H ( p ∗ ) (密度なら その 微分エントロピー)が 有限なら E [ log p θ ( X ) ] = − H ( p ∗ ) − KL ( p ∗ ∥ p θ ) E[\log p_\theta(X)] = -H(p^{\ast}) - \operatorname{KL}(p^{\ast} \Vert p_\theta) E [ log p θ ( X )] = − H ( p ∗ ) − KL ( p ∗ ∥ p θ ) である。したがって 対数尤度の 期待値の 最大化は KL ( p ∗ ∥ p θ ) \operatorname{KL}(p^{\ast} \Vert p_\theta) KL ( p ∗ ∥ p θ ) の 最小化と 同値で、 p ∗ = p θ 0 p^{\ast} = p_{\theta_0} p ∗ = p θ 0 なら θ 0 \theta_0 θ 0 は 最大点である。
証明. 1:同じ値の 項を まとめると 1 n ∑ i log p θ ( x i ) = ∑ x p ^ n ( x ) log p θ ( x ) = − H ( p ^ n , p θ ) \frac{1}{n}\sum_i\log p_\theta(x_i) = \sum_x\hat{p}_n(x)\log p_\theta(x) = -H(\hat{p}_n, p_\theta) n 1 ∑ i log p θ ( x i ) = ∑ x p ^ n ( x ) log p θ ( x ) = − H ( p ^ n , p θ ) 。p ^ n \hat{p}_n p ^ n は 高々 n n n 点に 確率を もつので H ( p ^ n ) ≤ log n H(\hat{p}_n) \leq \log n H ( p ^ n ) ≤ log n (系 7.3)は 有限で θ \theta θ に よらず、7.1 節の 分解が 使える。2 も 同様で、最後の 主張は 定理 7.2 に よる。 □ \square □
x 1 , … , x n x_1, \dots, x_n x 1 , … , x n が p ∗ p^{\ast} p ∗ からの i.i.d. 標本の 実現値で E [ ∣ log p θ ( X ) ∣ ] < ∞ E[\lvert \log p_\theta(X) \rvert] < \infty E [∣ log p θ ( X )∣] < ∞ なら、大数の 法則( 22-statistics 第1章 定理 1.27)より 1 n ℓ ( θ ) \frac{1}{n}\ell(\theta) n 1 ℓ ( θ ) は E [ log p θ ( X ) ] E[\log p_\theta(X)] E [ log p θ ( X )] に 確率収束するので、最尤推定は 真の 分布に KL ダイバージェンスの 意味で 最も 近い モデルを 標本から 探している( 22-statistics 第3章 定理 3.32 の 一致性の 証明の 概略に 現れる 不等式は KL ( p θ 0 ∥ p θ ) > 0 \operatorname{KL}(p_{\theta_0} \Vert p_\theta) > 0 KL ( p θ 0 ∥ p θ ) > 0 その ものである)。真の 分布が モデルに 含まれなくても、パラメータ 空間が コンパクトで、各 x x x で θ ↦ log p θ ( x ) \theta \mapsto \log p_\theta(x) θ ↦ log p θ ( x ) が 連続、 E [ sup θ ∣ log p θ ( X ) ∣ ] < ∞ E[\sup_\theta \lvert \log p_\theta(X) \rvert] < \infty E [ sup θ ∣ log p θ ( X )∣] < ∞ であり、E [ log p θ ( X ) ] E[\log p_\theta(X)] E [ log p θ ( X )] の 最大点( H ( p ∗ ) H(p^{\ast}) H ( p ∗ ) が 有限なら KL ( p ∗ ∥ p θ ) \operatorname{KL}(p^{\ast} \Vert p_\theta) KL ( p ∗ ∥ p θ ) の 最小点) θ ∗ \theta^{\ast} θ ∗ が ただ 一つなら、最尤推定量は θ ∗ \theta^{\ast} θ ∗ に 確率収束する(White, 1982 年。主張のみ。一様な 大数の 法則を 使って 示される)。連続分布でも、 − 1 n ℓ ( θ ) -\frac{1}{n}\ell(\theta) − n 1 ℓ ( θ ) は 経験分布で 期待値を とった 交差エントロピーである。第1章の 言葉では、密度推定は 対数損失 − log q ( x ) -\log q(x) − log q ( x ) の 経験リスク最小化で、そのリスク H ( p ∗ , q ) H(p^{\ast}, q) H ( p ∗ , q ) は(H ( p ∗ ) H(p^{\ast}) H ( p ∗ ) が 有限なら) q = p ∗ q = p^{\ast} q = p ∗ でだけ 最小に なる(系 7.3)。分類でも 同様である。
命題 7.8 (交差エントロピー損失)Y \mathcal{Y} Y を 有限集合、 p ∗ ( ⋅ ∣ x ) p^{\ast}(\cdot \mid x) p ∗ ( ⋅ ∣ x ) を X = x X = x X = x のもとでの Y Y Y の 真の 条件付き分布と する。各 x x x に Y \mathcal{Y} Y 上の 分布 q ( ⋅ ∣ x ) q(\cdot \mid x) q ( ⋅ ∣ x ) を 対応させる 予測の、損失 − log q ( y ∣ x ) -\log q(y \mid x) − log q ( y ∣ x ) に ついての リスクは E [ H ( p ∗ ( ⋅ ∣ X ) ) ] + E [ KL ( p ∗ ( ⋅ ∣ X ) ∥ q ( ⋅ ∣ X ) ) ] E[H(p^{\ast}(\cdot \mid X))] + E[\operatorname{KL}(p^{\ast}(\cdot \mid X) \Vert q(\cdot \mid X))] E [ H ( p ∗ ( ⋅ ∣ X ))] + E [ KL ( p ∗ ( ⋅ ∣ X ) ∥ q ( ⋅ ∣ X ))] であり、q ( ⋅ ∣ x ) = p ∗ ( ⋅ ∣ x ) q(\cdot \mid x) = p^{\ast}(\cdot \mid x) q ( ⋅ ∣ x ) = p ∗ ( ⋅ ∣ x ) と なる 予測は ベイズ最適である。
証明. X = x X = x X = x のもとでの 条件付きリスクは H ( p ∗ ( ⋅ ∣ x ) , q ( ⋅ ∣ x ) ) H(p^{\ast}(\cdot \mid x), q(\cdot \mid x)) H ( p ∗ ( ⋅ ∣ x ) , q ( ⋅ ∣ x )) で、H ( p ∗ ( ⋅ ∣ x ) ) ≤ log ∣ Y ∣ H(p^{\ast}(\cdot \mid x)) \leq \log\lvert \mathcal{Y} \rvert H ( p ∗ ( ⋅ ∣ x )) ≤ log ∣ Y ∣ は 有限だから 7.1 節の 分解が 使える。第1章の 命題 1.6・補題 1.8 と 系 7.3 から 従う。 □ \square □
第2章の 命題 2.15 は ∣ Y ∣ = 2 \lvert \mathcal{Y} \rvert = 2 ∣ Y ∣ = 2 の 場合である。0-1 損失(第1章の 定理 1.10)と 違い、交差エントロピー損失は 確率 その ものを 当てる ことを 求める。
例 7.9 (KL ダイバージェンスの 向き)最尤推定が 最小に する KL ( p ∗ ∥ p θ ) \operatorname{KL}(p^{\ast} \Vert p_\theta) KL ( p ∗ ∥ p θ ) は p ∗ p^{\ast} p ∗ で 期待値を とるので、データの ある ところで p θ ≈ 0 p_\theta \approx 0 p θ ≈ 0 だと 大きな 罰を 受ける。 p ∗ = 1 2 N ( − 3 , 1 ) + 1 2 N ( 3 , 1 ) p^{\ast} = \frac{1}{2}N(-3, 1) + \frac{1}{2}N(3, 1) p ∗ = 2 1 N ( − 3 , 1 ) + 2 1 N ( 3 , 1 ) を N ( m , s 2 ) N(m, s^2) N ( m , s 2 ) で 近似すると、 KL ( p ∗ ∥ N ( m , s 2 ) ) \operatorname{KL}(p^{\ast} \Vert N(m, s^2)) KL ( p ∗ ∥ N ( m , s 2 )) の 最小点は 平均と 分散を 合わせた N ( 0 , 10 ) N(0, 10) N ( 0 , 10 ) で(問題 7.3)、データの ほとんどない 0 0 0 の 付近に 密度の 山を おく。逆向きの KL ( N ( m , s 2 ) ∥ p ∗ ) \operatorname{KL}(N(m, s^2) \Vert p^{\ast}) KL ( N ( m , s 2 ) ∥ p ∗ ) は 近似する 側で 期待値を とるので p ∗ ≈ 0 p^{\ast} \approx 0 p ∗ ≈ 0 の ところを 避け、数値的に 最小化すると、最小点は N ( ± 2.98 , 1.05 ) N(\pm 2.98, 1.05) N ( ± 2.98 , 1.05 ) 付近(片方の 山だけ。値は 約 0.689 0.689 0.689 )に なる。逆向きの KL ダイバージェンスは 7.4 節の 変分推論に 現れる。
7.3 混合ガウスモデルと EM アルゴリズム
定義 7.10 (混合ガウスモデル, Gaussian mixture model)混合係数 π k > 0 \pi_k > 0 π k > 0 (∑ k = 1 K π k = 1 \sum_{k=1}^{K}\pi_k = 1 ∑ k = 1 K π k = 1 )、平均 μ k ∈ R d \mu_k \in \mathbb{R}^d μ k ∈ R d 、正定値な 共分散行列 Σ k \Sigma_k Σ k の 組を θ \theta θ とし、N d ( μ , Σ ) N_d(\mu, \Sigma) N d ( μ , Σ ) の 密度( 22-statistics 第1章 定理 1.22)を φ ( x ; μ , Σ ) \varphi(x; \mu, \Sigma) φ ( x ; μ , Σ ) と して、 p θ ( x ) = ∑ k = 1 K π k φ ( x ; μ k , Σ k ) p_\theta(x) = \sum_{k=1}^{K}\pi_k\varphi(x; \mu_k, \Sigma_k) p θ ( x ) = ∑ k = 1 K π k φ ( x ; μ k , Σ k ) を 密度と する 分布を 混合ガウス分布と いう。
確率 π k \pi_k π k で 成分の 番号 Z = k Z = k Z = k を 選んでから X ∼ N d ( μ k , Σ k ) X \sim N_d(\mu_k, \Sigma_k) X ∼ N d ( μ k , Σ k ) を 生成すると、 X X X の 分布は p θ p_\theta p θ 、( X , Z ) (X, Z) ( X , Z ) の 同時分布は p θ ( x , k ) = π k φ ( x ; μ k , Σ k ) p_\theta(x, k) = \pi_k\varphi(x; \mu_k, \Sigma_k) p θ ( x , k ) = π k φ ( x ; μ k , Σ k ) である。潜在変数 Z Z Z の 事後確率 γ k ( x ) = π k φ ( x ; μ k , Σ k ) / p θ ( x ) \gamma_k(x) = \pi_k\varphi(x; \mu_k, \Sigma_k)/p_\theta(x) γ k ( x ) = π k φ ( x ; μ k , Σ k ) / p θ ( x ) を 負担率 (responsibility) と いう。対数尤度 ℓ ( θ ) = ∑ i log ∑ k π k φ ( x i ; μ k , Σ k ) \ell(\theta) = \sum_i\log\sum_k\pi_k\varphi(x_i; \mu_k, \Sigma_k) ℓ ( θ ) = ∑ i log ∑ k π k φ ( x i ; μ k , Σ k ) は 閉じた 形では 最大化できないが、各 x i x_i x i の 成分の 番号が わかっていれば、成分ごとの 割合・平均・共分散行列を 計算するだけで よい。EM アルゴリズムは、わからない 番号を、現在の パラメータでの 事後分布で 平均して 埋める ことを 繰り返す。
一般に、観測 x x x と 有限集合に 値を とる 潜在変数 z z z の 同時分布 p θ ( x , z ) > 0 p_\theta(x, z) > 0 p θ ( x , z ) > 0 を もつ モデルを 考える( x , z x, z x , z は データ全体・潜在変数全体で よい)。 ℓ ( θ ) = log p θ ( x ) \ell(\theta) = \log p_\theta(x) ℓ ( θ ) = log p θ ( x ) 、p θ ( x ) = ∑ z p θ ( x , z ) p_\theta(x) = \sum_zp_\theta(x, z) p θ ( x ) = ∑ z p θ ( x , z ) である。
定義 7.11 (EM アルゴリズム, expectation–maximization algorithm)初期値 θ ( 0 ) \theta^{(0)} θ ( 0 ) から、t = 0 , 1 , 2 , … t = 0, 1, 2, \dots t = 0 , 1 , 2 , … に ついて 次を 繰り返す。
E ステップ :事後分布 p θ ( t ) ( z ∣ x ) p_{\theta^{(t)}}(z \mid x) p θ ( t ) ( z ∣ x ) を 求め、 Q ( θ ∣ θ ( t ) ) = ∑ z p θ ( t ) ( z ∣ x ) log p θ ( x , z ) Q(\theta \mid \theta^{(t)}) = \sum_zp_{\theta^{(t)}}(z \mid x)\log p_\theta(x, z) Q ( θ ∣ θ ( t ) ) = ∑ z p θ ( t ) ( z ∣ x ) log p θ ( x , z ) と おく。
M ステップ :Q ( θ ∣ θ ( t ) ) Q(\theta \mid \theta^{(t)}) Q ( θ ∣ θ ( t ) ) を 最大に する θ \theta θ を θ ( t + 1 ) \theta^{(t+1)} θ ( t + 1 ) と する。
データが i.i.d. なら 事後分布も ∏ i p θ ( z i ∣ x i ) \prod_ip_\theta(z_i \mid x_i) ∏ i p θ ( z i ∣ x i ) と 分かれ、 Q = ∑ i ∑ z i p θ ( t ) ( z i ∣ x i ) log p θ ( x i , z i ) Q = \sum_i\sum_{z_i}p_{\theta^{(t)}}(z_i \mid x_i)\log p_\theta(x_i, z_i) Q = ∑ i ∑ z i p θ ( t ) ( z i ∣ x i ) log p θ ( x i , z i ) と なる。
定理 7.12 (EM アルゴリズムの 単調性) Q ( θ ∣ θ ′ ) ≥ Q ( θ ′ ∣ θ ′ ) Q(\theta \mid \theta') \geq Q(\theta' \mid \theta') Q ( θ ∣ θ ′ ) ≥ Q ( θ ′ ∣ θ ′ ) ならば ℓ ( θ ) ≥ ℓ ( θ ′ ) \ell(\theta) \geq \ell(\theta') ℓ ( θ ) ≥ ℓ ( θ ′ ) である。特に EM アルゴリズムの 列に ついて ℓ ( θ ( 0 ) ) ≤ ℓ ( θ ( 1 ) ) ≤ ℓ ( θ ( 2 ) ) ≤ ⋯ \ell(\theta^{(0)}) \leq \ell(\theta^{(1)}) \leq \ell(\theta^{(2)}) \leq \cdots ℓ ( θ ( 0 ) ) ≤ ℓ ( θ ( 1 ) ) ≤ ℓ ( θ ( 2 ) ) ≤ ⋯ 。
証明. w ( z ) = p θ ′ ( z ∣ x ) w(z) = p_{\theta'}(z \mid x) w ( z ) = p θ ′ ( z ∣ x ) と おくと p θ ′ ( x , z ) = p θ ′ ( x ) w ( z ) p_{\theta'}(x, z) = p_{\theta'}(x)w(z) p θ ′ ( x , z ) = p θ ′ ( x ) w ( z ) なので、p θ ( x ) p θ ′ ( x ) = ∑ z w ( z ) p θ ( x , z ) p θ ′ ( x , z ) \frac{p_\theta(x)}{p_{\theta'}(x)} = \sum_zw(z)\frac{p_\theta(x, z)}{p_{\theta'}(x, z)} p θ ′ ( x ) p θ ( x ) = ∑ z w ( z ) p θ ′ ( x , z ) p θ ( x , z ) である。w w w は 確率分布で log \log log は 凹だから、イェンセンの 不等式( 01-calculus 第4章 定理 4.31 を 凸関数 − log -\log − log に 使う)より
ℓ ( θ ) − ℓ ( θ ′ ) ≥ ∑ z w ( z ) log p θ ( x , z ) p θ ′ ( x , z ) = Q ( θ ∣ θ ′ ) − Q ( θ ′ ∣ θ ′ ) ≥ 0 \ell(\theta) - \ell(\theta') \geq \sum_zw(z)\log\frac{p_\theta(x, z)}{p_{\theta'}(x, z)} = Q(\theta \mid \theta') - Q(\theta' \mid \theta') \geq 0 ℓ ( θ ) − ℓ ( θ ′ ) ≥ z ∑ w ( z ) log p θ ′ ( x , z ) p θ ( x , z ) = Q ( θ ∣ θ ′ ) − Q ( θ ′ ∣ θ ′ ) ≥ 0
M ステップの θ ( t + 1 ) \theta^{(t+1)} θ ( t + 1 ) は Q ( θ ( t + 1 ) ∣ θ ( t ) ) ≥ Q ( θ ( t ) ∣ θ ( t ) ) Q(\theta^{(t+1)} \mid \theta^{(t)}) \geq Q(\theta^{(t)} \mid \theta^{(t)}) Q ( θ ( t + 1 ) ∣ θ ( t ) ) ≥ Q ( θ ( t ) ∣ θ ( t ) ) を 満たす。 □ \square □
定理 7.13 (混合ガウスモデルの EM)データ x 1 , … , x n ∈ R d x_1, \dots, x_n \in \mathbb{R}^d x 1 , … , x n ∈ R d の すべてを 含む超 平面は ない( d = 1 d = 1 d = 1 なら、すべては 等しくない)と する。E ステップは 現在の パラメータでの 負担率 γ i k = γ k ( x i ) \gamma_{ik} = \gamma_k(x_i) γ ik = γ k ( x i ) の 計算であり、M ステップの 最大点は ただ 一つで、 N k = ∑ i γ i k N_k = \sum_i\gamma_{ik} N k = ∑ i γ ik と して
π k n e w = N k n , μ k n e w = 1 N k ∑ i = 1 n γ i k x i , Σ k n e w = 1 N k ∑ i = 1 n γ i k ( x i − μ k n e w ) ( x i − μ k n e w ) ⊤ \pi_k^{\mathrm{new}} = \frac{N_k}{n}, \qquad \mu_k^{\mathrm{new}} = \frac{1}{N_k}\sum_{i=1}^{n}\gamma_{ik}x_i, \qquad \Sigma_k^{\mathrm{new}} = \frac{1}{N_k}\sum_{i=1}^{n}\gamma_{ik}(x_i - \mu_k^{\mathrm{new}})(x_i - \mu_k^{\mathrm{new}})^{\top} π k new = n N k , μ k new = N k 1 i = 1 ∑ n γ ik x i , Σ k new = N k 1 i = 1 ∑ n γ ik ( x i − μ k new ) ( x i − μ k new ) ⊤
である。π k n e w > 0 \pi_k^{\mathrm{new}} > 0 π k new > 0 で、Σ k n e w \Sigma_k^{\mathrm{new}} Σ k new は 正定値である。
証明. P ( Z i = k ∣ x i ) = γ i k > 0 P(Z_i = k \mid x_i) = \gamma_{ik} > 0 P ( Z i = k ∣ x i ) = γ ik > 0 で、Q = ∑ k N k log π k + ∑ k ∑ i γ i k log φ ( x i ; μ k , Σ k ) Q = \sum_kN_k\log\pi_k + \sum_k\sum_i\gamma_{ik}\log\varphi(x_i; \mu_k, \Sigma_k) Q = ∑ k N k log π k + ∑ k ∑ i γ ik log φ ( x i ; μ k , Σ k ) は π \pi π の 部分と 各 ( μ k , Σ k ) (\mu_k, \Sigma_k) ( μ k , Σ k ) の 部分に 分かれる。 ∑ k γ i k = 1 \sum_k\gamma_{ik} = 1 ∑ k γ ik = 1 より ν k = N k / n \nu_k = N_k/n ν k = N k / n は 確率分布で、第 1 項は − n H ( ν , π ) -nH(\nu, \pi) − n H ( ν , π ) だから、系 7.3 より π = ν \pi = \nu π = ν でだけ 最大に なる。 x ˉ k = μ k n e w \bar{x}_k = \mu_k^{\mathrm{new}} x ˉ k = μ k new 、S k = Σ k n e w S_k = \Sigma_k^{\mathrm{new}} S k = Σ k new と おく。 ∑ i γ i k ( x i − x ˉ k ) = 0 \sum_i\gamma_{ik}(x_i - \bar{x}_k) = 0 ∑ i γ ik ( x i − x ˉ k ) = 0 で 交差項が 消える ことと v ⊤ Σ k − 1 v = tr ( Σ k − 1 v v ⊤ ) v^{\top}\Sigma_k^{-1}v = \operatorname{tr}(\Sigma_k^{-1}vv^{\top}) v ⊤ Σ k − 1 v = tr ( Σ k − 1 v v ⊤ ) から、第 2 項の k k k 番目は 定数を 除いて
− N k 2 ( log det Σ k + tr ( Σ k − 1 S k ) + ( x ˉ k − μ k ) ⊤ Σ k − 1 ( x ˉ k − μ k ) ) -\frac{N_k}{2}\Bigl(\log\det\Sigma_k + \operatorname{tr}(\Sigma_k^{-1}S_k) + (\bar{x}_k - \mu_k)^{\top}\Sigma_k^{-1}(\bar{x}_k - \mu_k)\Bigr) − 2 N k ( log det Σ k + tr ( Σ k − 1 S k ) + ( x ˉ k − μ k ) ⊤ Σ k − 1 ( x ˉ k − μ k ) )
で、どの Σ k \Sigma_k Σ k に ついても μ k = x ˉ k \mu_k = \bar{x}_k μ k = x ˉ k でだけ 最大に なる。仮定より v ≠ 0 v \neq 0 v = 0 なら v ⊤ x i v^{\top}x_i v ⊤ x i は すべては 等しくないので v ⊤ S k v = 1 N k ∑ i γ i k ( v ⊤ ( x i − x ˉ k ) ) 2 > 0 v^{\top}S_kv = \frac{1}{N_k}\sum_i\gamma_{ik}(v^{\top}(x_i - \bar{x}_k))^2 > 0 v ⊤ S k v = N k 1 ∑ i γ ik ( v ⊤ ( x i − x ˉ k ) ) 2 > 0 、すな わち S k S_k S k は 正定値である。正の 平方根 S k 1 / 2 S_k^{1/2} S k 1/2 (02-linear-algebra 第8章 命題 8.15)に ついて S k 1 / 2 Σ k − 1 S k 1 / 2 S_k^{1/2}\Sigma_k^{-1}S_k^{1/2} S k 1/2 Σ k − 1 S k 1/2 の 固有値を λ j > 0 \lambda_j > 0 λ j > 0 と すると、 tr ( Σ k − 1 S k ) = ∑ j λ j \operatorname{tr}(\Sigma_k^{-1}S_k) = \sum_j\lambda_j tr ( Σ k − 1 S k ) = ∑ j λ j 、log det Σ k = log det S k − ∑ j log λ j \log\det\Sigma_k = \log\det S_k - \sum_j\log\lambda_j log det Σ k = log det S k − ∑ j log λ j だから、λ − log λ ≥ 1 \lambda - \log\lambda \geq 1 λ − log λ ≥ 1 (定理 7.2 の 証明)より log det Σ k + tr ( Σ k − 1 S k ) ≥ log det S k + d \log\det\Sigma_k + \operatorname{tr}(\Sigma_k^{-1}S_k) \geq \log\det S_k + d log det Σ k + tr ( Σ k − 1 S k ) ≥ log det S k + d で、等号は すべての λ j = 1 \lambda_j = 1 λ j = 1 、すな わち Σ k = S k \Sigma_k = S_k Σ k = S k の ときに 限る。 □ \square □
例 7.14 (1 次元の EM)データ 1 , 2 , 3 , 4 , 8 , 9 , 10 1, 2, 3, 4, 8, 9, 10 1 , 2 , 3 , 4 , 8 , 9 , 10 に K = 2 K = 2 K = 2 の 混合ガウス分布を 当てはめる。初期値 π = ( 0.5 , 0.5 ) \pi = (0.5, 0.5) π = ( 0.5 , 0.5 ) 、μ = ( 2 , 4 ) \mu = (2, 4) μ = ( 2 , 4 ) 、σ 2 = ( 1 , 1 ) \sigma^2 = (1, 1) σ 2 = ( 1 , 1 ) での 負担率は γ i 1 = 1 / ( 1 + e 2 x i − 6 ) = σ ( 6 − 2 x i ) \gamma_{i1} = 1/(1 + e^{2x_i - 6}) = \sigma(6 - 2x_i) γ i 1 = 1/ ( 1 + e 2 x i − 6 ) = σ ( 6 − 2 x i ) (第2章の ロジスティック関数)で、 x i = 1 , 2 , 3 , 4 x_i = 1, 2, 3, 4 x i = 1 , 2 , 3 , 4 に 対して 0.982 , 0.881 , 0.5 , 0.119 0.982, 0.881, 0.5, 0.119 0.982 , 0.881 , 0.5 , 0.119 、8 , 9 , 10 8, 9, 10 8 , 9 , 10 に 対しては 10 − 4 10^{-4} 1 0 − 4 未満である。更新を t t t 回 行った 後の 値は 次の とおりである(Python で 計算した)。
t t t
ℓ \ell ℓ
π 1 \pi_1 π 1
μ 1 \mu_1 μ 1
μ 2 \mu_2 μ 2
σ 1 2 \sigma_1^2 σ 1 2
σ 2 2 \sigma_2^2 σ 2 2
0
− 49.8194 -49.8194 − 49.8194
0.5000 0.5000 0.5000
2.0000 2.0000 2.0000
4.0000 4.0000 4.0000
1.0000 1.0000 1.0000
1.0000 1.0000 1.0000
1
− 17.0402 -17.0402 − 17.0402
0.3546 0.3546 0.3546
1.9020 1.9020 1.9020
7.1447 7.1447 7.1447
0.7804 0.7804 0.7804
7.4061 7.4061 7.4061
2
− 16.8994 -16.8994 − 16.8994
0.3879 0.3879 0.3879
2.0422 2.0422 2.0422
7.3412 7.3412 7.3412
0.8448 0.8448 0.8448
7.1101 7.1101 7.1101
4
− 16.4717 -16.4717 − 16.4717
0.4658 0.4658 0.4658
2.2579 2.2579 2.2579
7.9255 7.9255 7.9255
1.0707 1.0707 1.0707
5.3448 5.3448 5.3448
6
− 15.1382 -15.1382 − 15.1382
0.5514 0.5514 0.5514
2.4513 2.4513 2.4513
8.7697 8.7697 8.7697
1.2200 1.2200 1.2200
1.7811 1.7811 1.7811
8
− 14.5510 -14.5510 − 14.5510
0.5714 0.5714 0.5714
2.5000 2.5000 2.5000
9.0000 9.0000 9.0000
1.2501 1.2501 1.2501
0.6667 0.6667 0.6667
対数尤度は 単調に 増え、 { 1 , 2 , 3 , 4 } \lbrace 1, 2, 3, 4 \rbrace { 1 , 2 , 3 , 4 } と { 8 , 9 , 10 } \lbrace 8, 9, 10 \rbrace { 8 , 9 , 10 } に 分けた ときの 割合 4 / 7 4/7 4/7 、平均 2.5 , 9 2.5, 9 2.5 , 9 、分散 1.25 , 2 / 3 1.25, 2/3 1.25 , 2/3 に ほぼ等しい 値に 落ち着く。
EM が 保証するのは、対数尤度が 減らない ことだけである。 ∇ θ Q ( θ ∣ θ ′ ) \nabla_\theta Q(\theta \mid \theta') ∇ θ Q ( θ ∣ θ ′ ) を θ = θ ′ \theta = \theta' θ = θ ′ で 評価すると ∑ z p θ ′ ( z ∣ x ) ∇ log p θ ′ ( x , z ) = ∇ ℓ ( θ ′ ) \sum_zp_{\theta'}(z \mid x)\nabla\log p_{\theta'}(x, z) = \nabla\ell(\theta') ∑ z p θ ′ ( z ∣ x ) ∇ log p θ ′ ( x , z ) = ∇ ℓ ( θ ′ ) なので、EM が 止まる 点は( Q Q Q の 最大点が 内部に あれば) ℓ \ell ℓ の 停留点である。さらに、上位集合 { θ ∣ ℓ ( θ ) ≥ ℓ ( θ ( 0 ) ) } \lbrace \theta \mid \ell(\theta) \geq \ell(\theta^{(0)}) \rbrace { θ ∣ ℓ ( θ ) ≥ ℓ ( θ ( 0 ) )} が コンパクトで、 ℓ \ell ℓ が 連続かつ内部で 微分可能、列が パラメータ 空間の 内部にと どまり、 Q ( θ ∣ θ ′ ) Q(\theta \mid \theta') Q ( θ ∣ θ ′ ) が ( θ , θ ′ ) (\theta, \theta') ( θ , θ ′ ) に ついて 連続ならば、EM の 列の 集積点は すべて ℓ \ell ℓ の 停留点で、 ℓ ( θ ( t ) ) \ell(\theta^{(t)}) ℓ ( θ ( t ) ) は ある 停留点での 値に 単調に 収束する(Wu, 1983 年。主張のみ。混合ガウスモデルでは、次の 命題 7.15 の ために この 上位集合が コンパクトでなく、そのままでは 使えない)。しかし 停留点は 最大点とは 限らない。例 7.14 で μ 1 = μ 2 \mu_1 = \mu_2 μ 1 = μ 2 、σ 1 2 = σ 2 2 \sigma_1^2 = \sigma_2^2 σ 1 2 = σ 2 2 から 始めると、すべての 負担率が π k \pi_k π k に 等しくなり、両成分とも 1 つの 正規分布の 最尤推定値 μ k = 37 / 7 \mu_k = 37/7 μ k = 37/7 、σ k 2 = 556 / 49 \sigma_k^2 = 556/49 σ k 2 = 556/49 に 移って 以後 動かない( ℓ ≈ − 18.4339 < − 14.5510 \ell \approx -18.4339 < -14.5510 ℓ ≈ − 18.4339 < − 14.5510 )。そもそも 尤度には 最大値が ない。
命題 7.15 (尤度の 非有界性) d = 1 d = 1 d = 1 、K = 2 K = 2 K = 2 とし、データ x 1 , … , x n x_1, \dots, x_n x 1 , … , x n を 任意にとる。 π 1 = π 2 = 1 / 2 \pi_1 = \pi_2 = 1/2 π 1 = π 2 = 1/2 、μ 1 = x 1 \mu_1 = x_1 μ 1 = x 1 、μ 2 = 0 \mu_2 = 0 μ 2 = 0 、σ 2 2 = 1 \sigma_2^2 = 1 σ 2 2 = 1 と 固定して σ 1 → + 0 \sigma_1 \to +0 σ 1 → + 0 と すると ℓ ( θ ) → ∞ \ell(\theta) \to \infty ℓ ( θ ) → ∞ 。したがって 混合ガウスモデルの 尤度は 上に 有界でなく、最尤推定量は 存在しない。
証明. p θ ( x 1 ) ≥ 1 2 φ ( x 1 ; x 1 , σ 1 2 ) = 1 2 2 π σ 1 p_\theta(x_1) \geq \frac{1}{2}\varphi(x_1; x_1, \sigma_1^2) = \frac{1}{2\sqrt{2\pi}\sigma_1} p θ ( x 1 ) ≥ 2 1 φ ( x 1 ; x 1 , σ 1 2 ) = 2 2 π σ 1 1 で、p θ ( x i ) ≥ 1 2 φ ( x i ; 0 , 1 ) p_\theta(x_i) \geq \frac{1}{2}\varphi(x_i; 0, 1) p θ ( x i ) ≥ 2 1 φ ( x i ; 0 , 1 ) は σ 1 \sigma_1 σ 1 に よらない 正の 数である。よって ℓ ( θ ) ≥ − log ( 2 2 π σ 1 ) + ∑ i ≥ 2 log φ ( x i ; 0 , 1 ) 2 → ∞ \ell(\theta) \geq -\log(2\sqrt{2\pi}\sigma_1) + \sum_{i \geq 2}\log\frac{\varphi(x_i; 0, 1)}{2} \to \infty ℓ ( θ ) ≥ − log ( 2 2 π σ 1 ) + ∑ i ≥ 2 log 2 φ ( x i ; 0 , 1 ) → ∞ 。□ \square □
成分が 1 点に 縮退すると、その 点の 密度が いくらでも 大きくなるのである(1 つの 正規分布では、観測値が すべて 等しい 場合を 除いて こうならない。 22-statistics 第3章 例 3.13)。例 7.14 の データで 初期値を π = ( 0.5 , 0.5 ) \pi = (0.5, 0.5) π = ( 0.5 , 0.5 ) 、μ = ( 1 , 6 ) \mu = (1, 6) μ = ( 1 , 6 ) 、σ 2 = ( 0.01 , 4 ) \sigma^2 = (0.01, 4) σ 2 = ( 0.01 , 4 ) と すると、1 回の 更新で σ 1 2 ≈ 2.9 × 10 − 20 \sigma_1^2 \approx 2.9 \times 10^{-20} σ 1 2 ≈ 2.9 × 1 0 − 20 、ℓ ≈ 3.39 \ell \approx 3.39 ℓ ≈ 3.39 と なり、次の 更新で σ 1 2 \sigma_1^2 σ 1 2 が 倍精度の 計算で 0 0 0 に なって 破綻する。成分の 番号を 入れ替えても p θ p_\theta p θ は 変わらないので、識別 可能性も 崩れている( 22-statistics 第3章 注意 3.34)。な お Σ k = ε I \Sigma_k = \varepsilon I Σ k = ε I 、π k = 1 / K \pi_k = 1/K π k = 1/ K と 固定すると γ i k \gamma_{ik} γ ik は e − ∥ x i − μ k ∥ 2 / ( 2 ε ) e^{-\lVert x_i - \mu_k \rVert^2/(2\varepsilon)} e − ∥ x i − μ k ∥ 2 / ( 2 ε ) に 比例するので、 ε → + 0 \varepsilon \to +0 ε → + 0 で 最も 近い 中心の 成分で 1 1 1 、ほかで 0 0 0 に 近づき(最も 近い 中心が ただ 一つの とき)、EM は 第3章 の k k k 平均法(定義 3.18)に 近づく。EM は k k k 平均法の 割り 当てを 確率に「柔らかく」した ものである。
ヒント
実務では
EM が 収束しても、最尤推定値が 得られたとは 限らない。(1) 初期値を 変えて 何回も 実行し( k k k 平均法の 結果を 初期値に する ことも 多い)、対数尤度が 最大の ものを 採る。(2) 分散に 下限を 設けるなどして 縮退を 防ぐ。同じ値が 繰り返し現れる データ(定価の 商品の 購入額、丸めた 測定値)では 特に 起こりやすく、異常に 大きい 対数尤度は まず 縮退を 疑う。(3) 成分の 数 K K K は 訓練データの 対数尤度では 選べない(成分を 複製すれば K + 1 K + 1 K + 1 成分で 同じ 分布を 表せるので、達成できる 値は K K K に ついて 下がらない)。検証データでの 対数尤度や BIC( 22-statistics 第6章 定義 6.16。最大対数尤度は 存在しないので、分散の 下限などの 制約のもとで EM が 見つけた 値で 代用する)で 選ぶ。
7.4 変分推論と ELBO
ベイズ統計(22-statistics 第7章 )では 未知の 量を まとめて z z z とし、事後分布 p ( z ∣ x ) = p ( x , z ) / p ( x ) p(z \mid x) = p(x, z)/p(x) p ( z ∣ x ) = p ( x , z ) / p ( x ) を 求めるが、周辺尤度 p ( x ) = ∫ p ( x , z ) d z p(x) = \int p(x, z)\ dz p ( x ) = ∫ p ( x , z ) d z は たいてい計算できない。 変分推論 (variational inference) は、扱いやすい 分布の 族 Q \mathcal{Q} Q から KL ( q ∥ p ( ⋅ ∣ x ) ) \operatorname{KL}(q \Vert p(\cdot \mid x)) KL ( q ∥ p ( ⋅ ∣ x )) が 最小の q q q を 探す。この 量は 未知の p ( x ) p(x) p ( x ) を 含むが、次の 恒等式で 避けられる。
定理 7.16 (ELBO 分解)p ( x ) > 0 p(x) > 0 p ( x ) > 0 とし、z z z の 分布 q q q は「q ( z ) > 0 q(z) > 0 q ( z ) > 0 なら p ( x , z ) > 0 p(x, z) > 0 p ( x , z ) > 0 」と E q [ ∣ log q ( Z ) ∣ ] < ∞ E_q[\lvert \log q(Z) \rvert] < \infty E q [∣ log q ( Z )∣] < ∞ を 満たすと する。 ELBO ( q ) = E q [ log p ( x , Z ) ] − E q [ log q ( Z ) ] \operatorname{ELBO}(q) = E_q[\log p(x, Z)] - E_q[\log q(Z)] ELBO ( q ) = E q [ log p ( x , Z )] − E q [ log q ( Z )] と おくと
log p ( x ) = ELBO ( q ) + KL ( q ∥ p ( ⋅ ∣ x ) ) \log p(x) = \operatorname{ELBO}(q) + \operatorname{KL}\bigl(q \Vert p(\cdot \mid x)\bigr) log p ( x ) = ELBO ( q ) + KL ( q ∥ p ( ⋅ ∣ x ) )
である(KL ダイバージェンスが ∞ \infty ∞ なら ELBO ( q ) = − ∞ \operatorname{ELBO}(q) = -\infty ELBO ( q ) = − ∞ と 読む)。特に ELBO ( q ) ≤ log p ( x ) \operatorname{ELBO}(q) \leq \log p(x) ELBO ( q ) ≤ log p ( x ) で、等号は q = p ( ⋅ ∣ x ) q = p(\cdot \mid x) q = p ( ⋅ ∣ x ) の ときに 限る。
証明. q ( z ) > 0 q(z) > 0 q ( z ) > 0 なら p ( z ∣ x ) = p ( x , z ) / p ( x ) > 0 p(z \mid x) = p(x, z)/p(x) > 0 p ( z ∣ x ) = p ( x , z ) / p ( x ) > 0 で、log p ( x ) = ( log p ( x , z ) − log q ( z ) ) + ( log q ( z ) − log p ( z ∣ x ) ) \log p(x) = (\log p(x, z) - \log q(z)) + (\log q(z) - \log p(z \mid x)) log p ( x ) = ( log p ( x , z ) − log q ( z )) + ( log q ( z ) − log p ( z ∣ x )) である。Z ∼ q Z \sim q Z ∼ q に ついて 期待値を とると、右辺の 第 2 項は KL ( q ∥ p ( ⋅ ∣ x ) ) ∈ [ 0 , ∞ ] \operatorname{KL}(q \Vert p(\cdot \mid x)) \in [0, \infty] KL ( q ∥ p ( ⋅ ∣ x )) ∈ [ 0 , ∞ ] 、第 1 項は ELBO ( q ) \operatorname{ELBO}(q) ELBO ( q ) に なる。後半は 定理 7.2 に よる。 □ \square □
p ( x ) p(x) p ( x ) を 証拠 (evidence) とも いうので、 ELBO ( q ) \operatorname{ELBO}(q) ELBO ( q ) を 証拠下界 (evidence lower bound) と いう。ELBO の 最大化は KL ダイバージェンスの 最小化と 同値で、計算には p ( x , z ) p(x, z) p ( x , z ) しか 要らない。 z z z の 事前分布 p Z p_Z p Z を 使えば ELBO ( q ) = E q [ log p ( x ∣ Z ) ] − KL ( q ∥ p Z ) \operatorname{ELBO}(q) = E_q[\log p(x \mid Z)] - \operatorname{KL}(q \Vert p_Z) ELBO ( q ) = E q [ log p ( x ∣ Z )] − KL ( q ∥ p Z ) (データへの 当ては まりと、事前分布から 離れる ことへの 罰)とも 書ける。最尤推定とは 逆向きの KL ダイバージェンスなので、 q q q は 例 7.9 のように 事後分布の 確率が 小さい ところを 避ける。
定理 7.17 (平均場近似の 座標上昇) z = ( z 1 , … , z m ) z = (z_1, \dots, z_m) z = ( z 1 , … , z m ) とし、q ( z ) = ∏ j q j ( z j ) q(z) = \prod_jq_j(z_j) q ( z ) = ∏ j q j ( z j ) の 形の 分布だけを 考える( 平均場近似 , mean-field approximation)。j j j 以外の q i q_i q i を 固定した とき、 ELBO ( q ) \operatorname{ELBO}(q) ELBO ( q ) を 最大に する q j q_j q j は ただ 一つで
q j ∗ ( z j ) = 1 C j exp ( E − j [ log p ( x , z j , Z − j ) ] ) q_j^{\ast}(z_j) = \frac{1}{C_j}\exp\Bigl(E_{-j}\bigl[\log p(x, z_j, Z_{-j})\bigr]\Bigr) q j ∗ ( z j ) = C j 1 exp ( E − j [ log p ( x , z j , Z − j ) ] )
である。ここで E − j E_{-j} E − j は i ≠ j i \neq j i = j の Z i ∼ q i Z_i \sim q_i Z i ∼ q i (互いに 独立)に ついての 期待値、 C j C_j C j は 正規化定数で、これらは すべて 有限と する。
証明. q q q は 積の 形なので E q [ log q ( Z ) ] = ∑ i E q i [ log q i ( Z i ) ] E_q[\log q(Z)] = \sum_iE_{q_i}[\log q_i(Z_i)] E q [ log q ( Z )] = ∑ i E q i [ log q i ( Z i )] であり、E q [ log p ( x , Z ) ] E_q[\log p(x, Z)] E q [ log p ( x , Z )] は 先に Z − j Z_{-j} Z − j に ついて 期待値を とれば E q j [ log q j ∗ ( Z j ) ] + log C j E_{q_j}[\log q_j^{\ast}(Z_j)] + \log C_j E q j [ log q j ∗ ( Z j )] + log C j に 等しい。よって q j q_j q j に よらない 定数 c j c_j c j に ついて ELBO ( q ) = E q j [ log q j ∗ ( Z j ) − log q j ( Z j ) ] + c j = − KL ( q j ∥ q j ∗ ) + c j \operatorname{ELBO}(q) = E_{q_j}[\log q_j^{\ast}(Z_j) - \log q_j(Z_j)] + c_j = -\operatorname{KL}(q_j \Vert q_j^{\ast}) + c_j ELBO ( q ) = E q j [ log q j ∗ ( Z j ) − log q j ( Z j )] + c j = − KL ( q j ∥ q j ∗ ) + c j で、定理 7.2 から 従う。 □ \square □
j = 1 , … , m j = 1, \dots, m j = 1 , … , m の 順に この 更新を 繰り返すと( 座標上昇変分推論 , CAVI)ELBO は 減らず、 log p ( x ) \log p(x) log p ( x ) 以下なので 値は 収束する。ただし ELBO は 一般に 凹でなく、結果は 初期値に 依存しうる。
例 7.18 (相関の ある 正規分布)事後 分布が N 2 ( μ , Σ ) N_2(\mu, \Sigma) N 2 ( μ , Σ ) 、Λ = Σ − 1 \Lambda = \Sigma^{-1} Λ = Σ − 1 の とき、 log p ( x , z ) = − 1 2 ( z − μ ) ⊤ Λ ( z − μ ) + c \log p(x, z) = -\frac{1}{2}(z - \mu)^{\top}\Lambda(z - \mu) + c log p ( x , z ) = − 2 1 ( z − μ ) ⊤ Λ ( z − μ ) + c の うち z 1 z_1 z 1 を 含む項は z 2 z_2 z 2 に ついて 1 次なので、定理 7.17 の 期待値は z 2 z_2 z 2 を m 2 = E q 2 [ Z 2 ] m_2 = E_{q_2}[Z_2] m 2 = E q 2 [ Z 2 ] で 置き換えて 平方完成すれば 求まり、 q 1 ∗ = N ( m 1 , 1 / Λ 11 ) q_1^{\ast} = N(m_1, 1/\Lambda_{11}) q 1 ∗ = N ( m 1 , 1/ Λ 11 ) 、m 1 = μ 1 − Λ 12 ( m 2 − μ 2 ) / Λ 11 m_1 = \mu_1 - \Lambda_{12}(m_2 - \mu_2)/\Lambda_{11} m 1 = μ 1 − Λ 12 ( m 2 − μ 2 ) / Λ 11 と なる( q 2 q_2 q 2 も 同様)。相関係数を ρ \rho ρ と すると、平均は μ \mu μ に 収束するが(誤差は 1 巡ごとに ρ 2 \rho^2 ρ 2 倍に なる)、分散 1 / Λ 11 = ( 1 − ρ 2 ) Σ 11 1/\Lambda_{11} = (1 - \rho^2)\Sigma_{11} 1/ Λ 11 = ( 1 − ρ 2 ) Σ 11 は 本当の 周辺分布の 分散 Σ 11 \Sigma_{11} Σ 11 より 小さく、 ρ = 0.9 \rho = 0.9 ρ = 0.9 なら 0.19 0.19 0.19 倍である(計算機でも 確かめた)。
注意
変分推論で 得た q q q から 作った 信用区間は、例 7.18 のように 狭すぎる ことがある。平均場近似は 変数の 間の 相関を 表せず、逆向きの KL ダイバージェンスは 事後分布の 確率が 小さい ところを 避けるからである。不確かさの 評価が 重要なら、マルコフ連鎖モンテカルロ法の 結果と 比べて 確かめる。
命題 7.19 (EM は ELBO の 座標上昇である) ELBO ( q , θ ) = E q [ log p θ ( x , Z ) ] − E q [ log q ( Z ) ] \operatorname{ELBO}(q, \theta) = E_q[\log p_\theta(x, Z)] - E_q[\log q(Z)] ELBO ( q , θ ) = E q [ log p θ ( x , Z )] − E q [ log q ( Z )] と おく。
θ \theta θ を 固定すると、 q q q に ついての 最大点は q = p θ ( ⋅ ∣ x ) q = p_\theta(\cdot \mid x) q = p θ ( ⋅ ∣ x ) で、最大値は ℓ ( θ ) \ell(\theta) ℓ ( θ ) である。
q = p θ ( t ) ( ⋅ ∣ x ) q = p_{\theta^{(t)}}(\cdot \mid x) q = p θ ( t ) ( ⋅ ∣ x ) を 固定すると、 ELBO ( q , θ ) = Q ( θ ∣ θ ( t ) ) − E q [ log q ( Z ) ] \operatorname{ELBO}(q, \theta) = Q(\theta \mid \theta^{(t)}) - E_q[\log q(Z)] ELBO ( q , θ ) = Q ( θ ∣ θ ( t ) ) − E q [ log q ( Z )] である。
したがって E ステップは q q q に ついて、M ステップは θ \theta θ に ついての ELBO の 最大化であり、 ℓ ( θ ( t + 1 ) ) ≥ ELBO ( q , θ ( t + 1 ) ) ≥ ELBO ( q , θ ( t ) ) = ℓ ( θ ( t ) ) \ell(\theta^{(t+1)}) \geq \operatorname{ELBO}(q, \theta^{(t+1)}) \geq \operatorname{ELBO}(q, \theta^{(t)}) = \ell(\theta^{(t)}) ℓ ( θ ( t + 1 ) ) ≥ ELBO ( q , θ ( t + 1 ) ) ≥ ELBO ( q , θ ( t ) ) = ℓ ( θ ( t ) ) と なる。
証明. 1 は 定理 7.16 を p θ p_\theta p θ に 使えば よく、2 は 定義 その ものである。最後の 不等式は 順に 1、M ステップ、1 に よる。 □ \square □
つまり EM は、潜在変数の 分布 q q q に 制約を 置かず(E ステップで 事後分布を 厳密に 求め)、 θ \theta θ は 点推定する 変分推論である。事後 分布が 計算できなければ q q q を 平均場近似などに 制限し( 変分 EM )、θ \theta θ にも 事前分布を 置いて z z z に 含めれば、ベイズ推論と しての 変分推論に なる。
7.5 隠れマルコフモデル
設備の 状態(正常・異常)は 直接は 見えず、アラームの 有無だけが 記録される。混合モデルの 成分の 番号を、時間とともに マルコフ連鎖で 変化させたのが 隠れマルコフモデルである(音声や 品詞の 推定などにも 使われる)。
定義 7.20 (隠れマルコフモデル, hidden Markov model)状態 { 1 , … , K } \lbrace 1, \dots, K \rbrace { 1 , … , K } の 初期分布 π \pi π 、推移確率 a j k ≥ 0 a_{jk} \geq 0 a j k ≥ 0 (∑ k a j k = 1 \sum_ka_{jk} = 1 ∑ k a j k = 1 )、各状態 k k k の 出力分布 b k b_k b k (確率関数または 密度)を 与え、状態の 列 z = ( z 1 , … , z T ) z = (z_1, \dots, z_T) z = ( z 1 , … , z T ) と 観測の 列 x = ( x 1 , … , x T ) x = (x_1, \dots, x_T) x = ( x 1 , … , x T ) の 同時分布を
p ( x , z ) = π z 1 b z 1 ( x 1 ) ∏ t = 2 T a z t − 1 z t b z t ( x t ) p(x, z) = \pi_{z_1}b_{z_1}(x_1)\prod_{t=2}^{T}a_{z_{t-1}z_t}b_{z_t}(x_t) p ( x , z ) = π z 1 b z 1 ( x 1 ) t = 2 ∏ T a z t − 1 z t b z t ( x t )
で 定める モデルを 隠れマルコフモデル(HMM)と いう。
状態の 列は マルコフ連鎖( 11-probability 第6章 定義 6.1)で、状態が 与えられると 観測は 独立であり、 x t x_t x t は z t z_t z t だけに 依存する。 x 1 : t = ( x 1 , … , x t ) x_{1:t} = (x_1, \dots, x_t) x 1 : t = ( x 1 , … , x t ) と 書く。最後の 因子 a z T − 1 z T b z T ( x T ) a_{z_{T-1}z_T}b_{z_T}(x_T) a z T − 1 z T b z T ( x T ) を x T , z T x_T, z_T x T , z T に ついて 和(積分)を とると 1 1 1 に なる ことを 繰り返せば、 t ≤ T t \leq T t ≤ T に ついて
p ( x 1 : t , z 1 : t ) = π z 1 b z 1 ( x 1 ) ∏ s = 2 t a z s − 1 z s b z s ( x s ) (2) p(x_{1:t}, z_{1:t}) = \pi_{z_1}b_{z_1}(x_1)\prod_{s=2}^{t}a_{z_{s-1}z_s}b_{z_s}(x_s) \tag{2} p ( x 1 : t , z 1 : t ) = π z 1 b z 1 ( x 1 ) s = 2 ∏ t a z s − 1 z s b z s ( x s ) ( 2 )
である。尤度 p ( x ) = ∑ z p ( x , z ) p(x) = \sum_zp(x, z) p ( x ) = ∑ z p ( x , z ) は K T K^T K T 項の 和で、 K = 10 K = 10 K = 10 , T = 100 T = 100 T = 100 なら 10 100 10^{100} 1 0 100 項に なる。
定理 7.21 (前向きアルゴリズム, forward algorithm)α t ( k ) = p ( x 1 : t , Z t = k ) \alpha_t(k) = p(x_{1:t}, Z_t = k) α t ( k ) = p ( x 1 : t , Z t = k ) ((2) を z 1 : t − 1 z_{1:t-1} z 1 : t − 1 に ついて 和を とり z t = k z_t = k z t = k とした もの)と おくと
α 1 ( k ) = π k b k ( x 1 ) , α t + 1 ( k ) = b k ( x t + 1 ) ∑ j = 1 K α t ( j ) a j k ( 1 ≤ t < T ) , p ( x ) = ∑ k = 1 K α T ( k ) \alpha_1(k) = \pi_kb_k(x_1), \qquad \alpha_{t+1}(k) = b_k(x_{t+1})\sum_{j=1}^{K}\alpha_t(j)a_{jk} \quad (1 \leq t < T), \qquad p(x) = \sum_{k=1}^{K}\alpha_T(k) α 1 ( k ) = π k b k ( x 1 ) , α t + 1 ( k ) = b k ( x t + 1 ) j = 1 ∑ K α t ( j ) a j k ( 1 ≤ t < T ) , p ( x ) = k = 1 ∑ K α T ( k )
であり、p ( x ) p(x) p ( x ) は O ( T K 2 ) O(TK^2) O ( T K 2 ) 回の 四則演算で 計算できる。
証明. (2) より p ( x 1 : t + 1 , z 1 : t + 1 ) = p ( x 1 : t , z 1 : t ) a z t z t + 1 b z t + 1 ( x t + 1 ) p(x_{1:t+1}, z_{1:t+1}) = p(x_{1:t}, z_{1:t})a_{z_tz_{t+1}}b_{z_{t+1}}(x_{t+1}) p ( x 1 : t + 1 , z 1 : t + 1 ) = p ( x 1 : t , z 1 : t ) a z t z t + 1 b z t + 1 ( x t + 1 ) で、z t + 1 = k z_{t+1} = k z t + 1 = k を 固定して z 1 : t − 1 z_{1:t-1} z 1 : t − 1 と z t = j z_t = j z t = j に ついて 和を とれば 漸化式を 得る。 α 1 \alpha_1 α 1 は (2) の t = 1 t = 1 t = 1 の 場合で、 ∑ k α T ( k ) = ∑ z p ( x , z ) = p ( x ) \sum_k\alpha_T(k) = \sum_zp(x, z) = p(x) ∑ k α T ( k ) = ∑ z p ( x , z ) = p ( x ) である。各 t t t で K K K 個の α t + 1 ( k ) \alpha_{t+1}(k) α t + 1 ( k ) を それぞれ K + 1 K + 1 K + 1 回の 乗算と K − 1 K - 1 K − 1 回の 加算で 求めるので 1 段は O ( K 2 ) O(K^2) O ( K 2 ) 、全体で O ( T K 2 ) O(TK^2) O ( T K 2 ) である。□ \square □
素朴な 和は、 K T K^T K T 個の 項の それぞれに 2 T − 1 2T - 1 2 T − 1 回の 乗算が 要るので O ( T K T ) O(TK^T) O ( T K T ) である。マルコフ性に より、先の 計算に 要る 過去の 情報が α t \alpha_t α t の K K K 個の 数に 集約されるのであり、動的計画法( 23-optimization 第7章 7.8 節)の 一例である。
例 7.22 (アラームの 記録)状態 1 を 正常、2 を 異常、観測 1 1 1 を アラームあり、 0 0 0 を なしとし、 π = ( 0.9 , 0.1 ) \pi = (0.9, 0.1) π = ( 0.9 , 0.1 ) 、a 11 = 0.9 a_{11} = 0.9 a 11 = 0.9 , a 12 = 0.1 a_{12} = 0.1 a 12 = 0.1 , a 21 = 0.3 a_{21} = 0.3 a 21 = 0.3 , a 22 = 0.7 a_{22} = 0.7 a 22 = 0.7 、b 1 ( 1 ) = 0.1 b_1(1) = 0.1 b 1 ( 1 ) = 0.1 , b 2 ( 1 ) = 0.6 b_2(1) = 0.6 b 2 ( 1 ) = 0.6 と する。観測 x = ( 0 , 1 , 1 ) x = (0, 1, 1) x = ( 0 , 1 , 1 ) では 次のようになる。
t t t
x t x_t x t
α t ( 1 ) \alpha_t(1) α t ( 1 )
α t ( 2 ) \alpha_t(2) α t ( 2 )
1
0
0.9 × 0.9 = 0.81 0.9 \times 0.9 = 0.81 0.9 × 0.9 = 0.81
0.1 × 0.4 = 0.04 0.1 \times 0.4 = 0.04 0.1 × 0.4 = 0.04
2
1
0.1 × ( 0.81 × 0.9 + 0.04 × 0.3 ) = 0.0741 0.1 \times (0.81 \times 0.9 + 0.04 \times 0.3) = 0.0741 0.1 × ( 0.81 × 0.9 + 0.04 × 0.3 ) = 0.0741
0.6 × ( 0.81 × 0.1 + 0.04 × 0.7 ) = 0.0654 0.6 \times (0.81 \times 0.1 + 0.04 \times 0.7) = 0.0654 0.6 × ( 0.81 × 0.1 + 0.04 × 0.7 ) = 0.0654
3
1
0.1 × ( 0.0741 × 0.9 + 0.0654 × 0.3 ) = 0.008631 0.1 \times (0.0741 \times 0.9 + 0.0654 \times 0.3) = 0.008631 0.1 × ( 0.0741 × 0.9 + 0.0654 × 0.3 ) = 0.008631
0.6 × ( 0.0741 × 0.1 + 0.0654 × 0.7 ) = 0.031914 0.6 \times (0.0741 \times 0.1 + 0.0654 \times 0.7) = 0.031914 0.6 × ( 0.0741 × 0.1 + 0.0654 × 0.7 ) = 0.031914
p ( x ) = 0.040545 p(x) = 0.040545 p ( x ) = 0.040545 で、8 8 8 通りの 状態の 列に ついて p ( x , z ) p(x, z) p ( x , z ) を 足した ものと 一致する(計算機で 確かめた)。正規化した α t ( k ) / ∑ j α t ( j ) = P ( Z t = k ∣ x 1 : t ) \alpha_t(k)/\sum_j\alpha_t(j) = P(Z_t = k \mid x_{1:t}) α t ( k ) / ∑ j α t ( j ) = P ( Z t = k ∣ x 1 : t ) (フィルタリング )に よると、異常の 確率は 時刻 2 で 0.0654 / 0.1395 ≈ 0.469 0.0654/0.1395 \approx 0.469 0.0654/0.1395 ≈ 0.469 、時刻 3 で 0.031914 / 0.040545 ≈ 0.787 0.031914/0.040545 \approx 0.787 0.031914/0.040545 ≈ 0.787 である。
長い 系列では α t \alpha_t α t が 指数的に 小さくなる。この モデルで 0 0 0 と 1 1 1 を 交互に 並べた 長さ 1000 の 観測列では log p ( x ) ≈ − 919.65 \log p(x) \approx -919.65 log p ( x ) ≈ − 919.65 、すな わち p ( x ) ≈ 10 − 399.4 p(x) \approx 10^{-399.4} p ( x ) ≈ 1 0 − 399.4 で、倍精度で 表せる 最小の 正の 数(約 4.9 × 10 − 324 4.9 \times 10^{-324} 4.9 × 1 0 − 324 )より 小さく、漸化式を そのまま 計算すると 途中( t = 810 t = 810 t = 810 )で α t \alpha_t α t が 0 0 0 に なる(計算機で 確かめた)。そこで 各段で α t \alpha_t α t を 正規化し、正規化定数の 対数を 足していく。 t t t 段目の 正規化定数は p ( x t ∣ x 1 : t − 1 ) p(x_t \mid x_{1:t-1}) p ( x t ∣ x 1 : t − 1 ) であり、7.6 節の 自己回帰分解 log p ( x ) = ∑ t log p ( x t ∣ x 1 : t − 1 ) \log p(x) = \sum_t\log p(x_t \mid x_{1:t-1}) log p ( x ) = ∑ t log p ( x t ∣ x 1 : t − 1 ) を 計算している ことになる。
ビタビ・アルゴリズム (Viterbi algorithm):最も 確からしい 状態の 列 arg max z p ( x , z ) \arg\max_zp(x, z) arg max z p ( x , z ) は、漸化式の 和を 最大値に 変えた δ 1 ( k ) = π k b k ( x 1 ) \delta_1(k) = \pi_kb_k(x_1) δ 1 ( k ) = π k b k ( x 1 ) 、δ t + 1 ( k ) = b k ( x t + 1 ) max j δ t ( j ) a j k \delta_{t+1}(k) = b_k(x_{t+1})\max_j\delta_t(j)a_{jk} δ t + 1 ( k ) = b k ( x t + 1 ) max j δ t ( j ) a j k を 計算し、最大を 与えた j j j を 記録して 最後から 逆に たどれば、 O ( T K 2 ) O(TK^2) O ( T K 2 ) で 求まる。例 7.22 では(正常, 異常, 異常)で、その 事後確率は 0.020412 / 0.040545 ≈ 0.503 0.020412/0.040545 \approx 0.503 0.020412/0.040545 ≈ 0.503 である。
バウム–ウェルチ・アルゴリズム (Baum–Welch algorithm):パラメータが 未知なら EM アルゴリズムで 推定する。E ステップに 要る P ( Z t = k ∣ x ) = α t ( k ) β t ( k ) / p ( x ) P(Z_t = k \mid x) = \alpha_t(k)\beta_t(k)/p(x) P ( Z t = k ∣ x ) = α t ( k ) β t ( k ) / p ( x ) などは、後ろ 向きの 量 β t ( k ) = p ( x t + 1 : T ∣ Z t = k ) \beta_t(k) = p(x_{t+1:T} \mid Z_t = k) β t ( k ) = p ( x t + 1 : T ∣ Z t = k ) を 同様の 漸化式で 求めれば 計算でき(問題 7.5)、M ステップでは 推移の 回数の 期待値の 比などで パラメータを 更新する(導出は 省略する)。定理 7.12 に より 尤度は 単調に 増加するが( a j k = 0 a_{jk} = 0 a j k = 0 などで p θ ( x , z ) = 0 p_\theta(x, z) = 0 p θ ( x , z ) = 0 と なる z z z が あっても、その 証明の 和を p θ ′ ( z ∣ x ) > 0 p_{\theta'}(z \mid x) > 0 p θ ′ ( z ∣ x ) > 0 の z z z に 限れば、最初の 等式が ≥ \geq ≥ に なるだけで 成り立つ)、局所解の 問題は 混合ガウスモデルと 同じである。
7.6 生成モデル(紹介)
代表的な 3 つの 生成モデルを、定式化と 学習の 目的関数に 絞って 本章の 言葉で 読む(ネットワークの 構造や 性能には 立ち入らない)。
変分オートエンコーダ
変分オートエンコーダ (variational autoencoder, VAE) は、z ∼ N m ( 0 , I ) z \sim N_m(0, I) z ∼ N m ( 0 , I ) と ニューラルネットワーク f θ f_\theta f θ に よる p θ ( x ∣ z ) = N d ( f θ ( z ) , σ 2 I ) p_\theta(x \mid z) = N_d(f_\theta(z), \sigma^2I) p θ ( x ∣ z ) = N d ( f θ ( z ) , σ 2 I ) などの モデル( デコーダ )を、事後分布の 近似 q ϕ ( z ∣ x ) = N m ( μ ϕ ( x ) , diag ( s ϕ ( x ) 2 ) ) q_\phi(z \mid x) = N_m(\mu_\phi(x), \operatorname{diag}(s_\phi(x)^2)) q ϕ ( z ∣ x ) = N m ( μ ϕ ( x ) , diag ( s ϕ ( x ) 2 )) (エンコーダ 。これも ネットワーク)とともに、ELBO
E q ϕ ( z ∣ x ) [ log p θ ( x ∣ Z ) ] − KL ( q ϕ ( ⋅ ∣ x ) ∥ N m ( 0 , I ) ) ≤ log p θ ( x ) E_{q_\phi(z \mid x)}\bigl[\log p_\theta(x \mid Z)\bigr] - \operatorname{KL}\bigl(q_\phi(\cdot \mid x) \Vert N_m(0, I)\bigr) \leq \log p_\theta(x) E q ϕ ( z ∣ x ) [ log p θ ( x ∣ Z ) ] − KL ( q ϕ ( ⋅ ∣ x ) ∥ N m ( 0 , I ) ) ≤ log p θ ( x )
の データに ついての 和を 最大に して 学習する(変分 EM と 同じ ELBO を、 q q q を ネットワークで 表し、E ステップと M ステップを 交互に 解く 代わりに θ \theta θ と ϕ \phi ϕ に ついて 同時に 勾配法で 最大化する)。第 1 項は 定数を 除いて − 1 2 σ 2 E ∥ x − f θ ( Z ) ∥ 2 -\frac{1}{2\sigma^2}E\lVert x - f_\theta(Z) \rVert^2 − 2 σ 2 1 E ∥ x − f θ ( Z ) ∥ 2 (x x x を z z z に 符号化して 復元した ときの 誤差)、KL ダイバージェンスの 項は (1) を 成分ごとに 足して 1 2 ∑ j ( μ j 2 + s j 2 − 1 − log s j 2 ) \frac{1}{2}\sum_j(\mu_j^2 + s_j^2 - 1 - \log s_j^2) 2 1 ∑ j ( μ j 2 + s j 2 − 1 − log s j 2 ) (μ = μ ϕ ( x ) \mu = \mu_\phi(x) μ = μ ϕ ( x ) 、s = s ϕ ( x ) s = s_\phi(x) s = s ϕ ( x ) )である。学習後は z ∼ N m ( 0 , I ) z \sim N_m(0, I) z ∼ N m ( 0 , I ) から f θ ( z ) f_\theta(z) f θ ( z ) を 作れば 新しい データが 得られる。期待値を とる 分布が ϕ \phi ϕ に よるので、勾配は 次の 書き換えで 計算する。
命題 7.23 (再パラメータ化, reparameterization)ε ∼ N m ( 0 , I ) \varepsilon \sim N_m(0, I) ε ∼ N m ( 0 , I ) 、s j > 0 s_j > 0 s j > 0 なら、成分ごとの 積 ⊙ \odot ⊙ に ついて μ + s ⊙ ε ∼ N m ( μ , diag ( s 2 ) ) \mu + s \odot \varepsilon \sim N_m(\mu, \operatorname{diag}(s^2)) μ + s ⊙ ε ∼ N m ( μ , diag ( s 2 )) であり、E Z ∼ N m ( μ , diag ( s 2 ) ) [ g ( Z ) ] = E [ g ( μ + s ⊙ ε ) ] E_{Z \sim N_m(\mu, \operatorname{diag}(s^2))}[g(Z)] = E[g(\mu + s \odot \varepsilon)] E Z ∼ N m ( μ , diag ( s 2 )) [ g ( Z )] = E [ g ( μ + s ⊙ ε )] である。g g g が C 1 C^1 C 1 級で 微分と 期待値の 順序を 交換できるなら
∂ ∂ μ j E [ g ( μ + s ⊙ ε ) ] = E [ ∂ g ∂ z j ( μ + s ⊙ ε ) ] , ∂ ∂ s j E [ g ( μ + s ⊙ ε ) ] = E [ ε j ∂ g ∂ z j ( μ + s ⊙ ε ) ] \frac{\partial}{\partial \mu_j}E[g(\mu + s \odot \varepsilon)] = E\left[\frac{\partial g}{\partial z_j}(\mu + s \odot \varepsilon)\right], \qquad \frac{\partial}{\partial s_j}E[g(\mu + s \odot \varepsilon)] = E\left[\varepsilon_j\frac{\partial g}{\partial z_j}(\mu + s \odot \varepsilon)\right] ∂ μ j ∂ E [ g ( μ + s ⊙ ε )] = E [ ∂ z j ∂ g ( μ + s ⊙ ε ) ] , ∂ s j ∂ E [ g ( μ + s ⊙ ε )] = E [ ε j ∂ z j ∂ g ( μ + s ⊙ ε ) ]
証明. μ + s ⊙ ε = μ + diag ( s ) ε \mu + s \odot \varepsilon = \mu + \operatorname{diag}(s)\varepsilon μ + s ⊙ ε = μ + diag ( s ) ε だから、前半は 22-statistics 第1章 定理 1.22 の 2 に よる。後半は 期待値の 中を 連鎖律で 微分すれば よい(交換の 十分条件は 06-measure-integration 第3章 定理 3.25)。□ \square □
右辺は ε \varepsilon ε を 生成すれば 不偏に 推定でき、 μ = μ ϕ ( x ) \mu = \mu_\phi(x) μ = μ ϕ ( x ) 、s = s ϕ ( x ) s = s_\phi(x) s = s ϕ ( x ) を 通して 誤差逆伝播法(第5章 定理 5.5)で ϕ \phi ϕ の 勾配が 得られる。
拡散モデル
拡散モデル (diffusion model) は、データ x 0 ∼ p ∗ x_0 \sim p^{\ast} x 0 ∼ p ∗ に 雑音を 少し ずつ 加える 前向き過程 x t = 1 − β t x t − 1 + β t ε t x_t = \sqrt{1 - \beta_t}x_{t-1} + \sqrt{\beta_t}\varepsilon_t x t = 1 − β t x t − 1 + β t ε t (t = 1 , … , T t = 1, \dots, T t = 1 , … , T 。0 < β t < 1 0 < \beta_t < 1 0 < β t < 1 は 定数、 ε t \varepsilon_t ε t は x 0 x_0 x 0 と 独立な N d ( 0 , I ) N_d(0, I) N d ( 0 , I ) の i.i.d.)と、それを 逆に たどる、学習する 逆過程 からなる。
命題 7.24 (前向き過程の 周辺分布) α ˉ t = ∏ s = 1 t ( 1 − β s ) \bar{\alpha}_t = \prod_{s=1}^{t}(1 - \beta_s) α ˉ t = ∏ s = 1 t ( 1 − β s ) と おくと、 x 0 x_0 x 0 を 与えたもとで x t ∼ N d ( α ˉ t x 0 , ( 1 − α ˉ t ) I ) x_t \sim N_d(\sqrt{\bar{\alpha}_t}x_0, (1 - \bar{\alpha}_t)I) x t ∼ N d ( α ˉ t x 0 , ( 1 − α ˉ t ) I ) である。すな わち x t x_t x t は、x 0 x_0 x 0 と 独立な ε ∼ N d ( 0 , I ) \varepsilon \sim N_d(0, I) ε ∼ N d ( 0 , I ) に よる α ˉ t x 0 + 1 − α ˉ t ε \sqrt{\bar{\alpha}_t}x_0 + \sqrt{1 - \bar{\alpha}_t}\varepsilon α ˉ t x 0 + 1 − α ˉ t ε と 同じ 分布に 従う。
証明. u t = x t − α ˉ t x 0 u_t = x_t - \sqrt{\bar{\alpha}_t}x_0 u t = x t − α ˉ t x 0 と おくと、 1 − β t α ˉ t − 1 = α ˉ t \sqrt{1 - \beta_t}\sqrt{\bar{\alpha}_{t-1}} = \sqrt{\bar{\alpha}_t} 1 − β t α ˉ t − 1 = α ˉ t より u 0 = 0 u_0 = 0 u 0 = 0 、u t = 1 − β t u t − 1 + β t ε t u_t = \sqrt{1 - \beta_t}u_{t-1} + \sqrt{\beta_t}\varepsilon_t u t = 1 − β t u t − 1 + β t ε t である。帰納法に より u t u_t u t は ε 1 , … , ε t \varepsilon_1, \dots, \varepsilon_t ε 1 , … , ε t の 1 次結合なので、x 0 x_0 x 0 と 独立で、平均 0 0 0 の 多変量正規分布に 従う( 22-statistics 第1章 定義 1.21・定理 1.22)。共分散行列を v t I v_tI v t I と すると、独立性から v t = ( 1 − β t ) v t − 1 + β t v_t = (1 - \beta_t)v_{t-1} + \beta_t v t = ( 1 − β t ) v t − 1 + β t 、v 0 = 0 v_0 = 0 v 0 = 0 で、v t = 1 − α ˉ t v_t = 1 - \bar{\alpha}_t v t = 1 − α ˉ t が これを 満たす。 □ \square □
α ˉ T ≈ 0 \bar{\alpha}_T \approx 0 α ˉ T ≈ 0 なら x T x_T x T は データに よらず ほぼ N d ( 0 , I ) N_d(0, I) N d ( 0 , I ) に 従う( β t = 0.02 \beta_t = 0.02 β t = 0.02 、T = 200 T = 200 T = 200 なら α ˉ T ≈ 0.018 \bar{\alpha}_T \approx 0.018 α ˉ T ≈ 0.018 )。逆過程は x T ∼ N d ( 0 , I ) x_T \sim N_d(0, I) x T ∼ N d ( 0 , I ) から 正規分布 p θ ( x t − 1 ∣ x t ) = N d ( m θ ( x t , t ) , σ t 2 I ) p_\theta(x_{t-1} \mid x_t) = N_d(m_\theta(x_t, t), \sigma_t^2I) p θ ( x t − 1 ∣ x t ) = N d ( m θ ( x t , t ) , σ t 2 I ) で 順に x 0 x_0 x 0 まで 生成する(以下、分散 σ t 2 \sigma_t^2 σ t 2 は 固定する)。 x 1 : T x_{1:T} x 1 : T を 潜在変数、前向き過程を 学習しない 変分分布と みると 定理 7.16 の 下界が 得られ、その 符号を 変えた ものは、 p θ ( x t − 1 ∣ x t ) p_\theta(x_{t-1} \mid x_t) p θ ( x t − 1 ∣ x t ) と、閉じた 形で 書ける 正規分布
q ( x t − 1 ∣ x t , x 0 ) = N d ( 1 1 − β t ( x t − β t 1 − α ˉ t ε ) , ( 1 − α ˉ t − 1 ) β t 1 − α ˉ t I ) ( 2 ≤ t ≤ T ) q(x_{t-1} \mid x_t, x_0) = N_d\Bigl(\frac{1}{\sqrt{1 - \beta_t}}\Bigl(x_t - \frac{\beta_t}{\sqrt{1 - \bar{\alpha}_t}}\varepsilon\Bigr), \frac{(1 - \bar{\alpha}_{t-1})\beta_t}{1 - \bar{\alpha}_t}I\Bigr) \qquad (2 \leq t \leq T) q ( x t − 1 ∣ x t , x 0 ) = N d ( 1 − β t 1 ( x t − 1 − α ˉ t β t ε ) , 1 − α ˉ t ( 1 − α ˉ t − 1 ) β t I ) ( 2 ≤ t ≤ T )
との KL ダイバージェンスの 和などに 分かれる。ここで ε = ( x t − α ˉ t x 0 ) / 1 − α ˉ t \varepsilon = (x_t - \sqrt{\bar{\alpha}_t}x_0)/\sqrt{1 - \bar{\alpha}_t} ε = ( x t − α ˉ t x 0 ) / 1 − α ˉ t であり、この 式は、 q ( x t ∣ x t − 1 ) q(x_t \mid x_{t-1}) q ( x t ∣ x t − 1 ) と 命題 7.24 の q ( x t − 1 ∣ x 0 ) q(x_{t-1} \mid x_0) q ( x t − 1 ∣ x 0 ) の 積を x t − 1 x_{t-1} x t − 1 に ついて 平方完成すれば 得られる。平均 m θ ( x t , t ) m_\theta(x_t, t) m θ ( x t , t ) を、この 式の 平均に 現れる ε \varepsilon ε を 雑音を 予測する ネットワーク ε θ ( x t , t ) \varepsilon_\theta(x_t, t) ε θ ( x t , t ) に 置き換えた ものにとると、(1) を 成分ごとに 使えば 各項は β t 2 2 σ t 2 ( 1 − β t ) ( 1 − α ˉ t ) E ∥ ε − ε θ ( α ˉ t x 0 + 1 − α ˉ t ε , t ) ∥ 2 \frac{\beta_t^2}{2\sigma_t^2(1 - \beta_t)(1 - \bar{\alpha}_t)}E\lVert \varepsilon - \varepsilon_\theta(\sqrt{\bar{\alpha}_t}x_0 + \sqrt{1 - \bar{\alpha}_t}\varepsilon, t) \rVert^2 2 σ t 2 ( 1 − β t ) ( 1 − α ˉ t ) β t 2 E ∥ ε − ε θ ( α ˉ t x 0 + 1 − α ˉ t ε , t ) ∥ 2 に θ \theta θ に よらない 定数を 加えた ものに なる( t = 1 t = 1 t = 1 の 項 − E [ log p θ ( x 0 ∣ x 1 ) ] -E[\log p_\theta(x_0 \mid x_1)] − E [ log p θ ( x 0 ∣ x 1 )] も 同じ形で、 KL ( q ( x T ∣ x 0 ) ∥ N d ( 0 , I ) ) \operatorname{KL}(q(x_T \mid x_0) \Vert N_d(0, I)) KL ( q ( x T ∣ x 0 ) ∥ N d ( 0 , I )) は θ \theta θ に よらない)。実際には、 t t t を 一様に 選び、この 係数を 省いた 二乗誤差を 最小に する ことが 多い(係数を 省くのは 経験的な 選択である)。第1章の 定理 1.9 より 最適な 予測は E [ ε ∣ x t ] E[\varepsilon \mid x_t] E [ ε ∣ x t ] で、これは x t x_t x t の 密度 q t q_t q t の 対数の 勾配( スコア )を 使って − 1 − α ˉ t ∇ log q t ( x t ) -\sqrt{1 - \bar{\alpha}_t}\nabla\log q_t(x_t) − 1 − α ˉ t ∇ log q t ( x t ) と 表せる(問題 7.6)。刻み幅 h h h に ついて β t = β ( t h ) h \beta_t = \beta(th)h β t = β ( t h ) h と おいて h → 0 h \to 0 h → 0 と する 極限では、前向き過程は 確率微分方程式 d X s = − 1 2 β ( s ) X s d s + β ( s ) d B s dX_s = -\frac{1}{2}\beta(s)X_s\ ds + \sqrt{\beta(s)}\ dB_s d X s = − 2 1 β ( s ) X s d s + β ( s ) d B s の 解に 近づく(主張のみ)。 β \beta β が 定数なら、各成分は 11-probability 第7章 例 7.20 の オルンシュタイン–ウーレンベック過程で、同例より、出発点に よらず 分散 ( β ) 2 / ( 2 ⋅ β / 2 ) = 1 (\sqrt{\beta})^2/(2 \cdot \beta/2) = 1 ( β ) 2 / ( 2 ⋅ β /2 ) = 1 の 正規分布 N ( 0 , 1 ) N(0, 1) N ( 0 , 1 ) に 近づく。
大規模言語モデル
文章を、有限集合 V V V (語彙 )の 元である トークン(単語や その 断片)の 列 x 1 , … , x T x_1, \dots, x_T x 1 , … , x T で 表し、 x < t = ( x 1 , … , x t − 1 ) x_{< t} = (x_1, \dots, x_{t-1}) x < t = ( x 1 , … , x t − 1 ) と 書く。
命題 7.25 (自己回帰分解)V T V^T V T 上の 任意の 確率分布 p p p と、p ( x 1 , … , x T ) > 0 p(x_1, \dots, x_T) > 0 p ( x 1 , … , x T ) > 0 と なる 列に ついて
p ( x 1 , … , x T ) = ∏ t = 1 T p ( x t ∣ x < t ) p(x_1, \dots, x_T) = \prod_{t=1}^{T}p(x_t \mid x_{< t}) p ( x 1 , … , x T ) = t = 1 ∏ T p ( x t ∣ x < t )
が 成り立つ。ここで p ( x t ∣ x < t ) = p ( x ≤ t ) / p ( x < t ) p(x_t \mid x_{< t}) = p(x_{\leq t})/p(x_{< t}) p ( x t ∣ x < t ) = p ( x ≤ t ) / p ( x < t ) は 最初の t t t 個と t − 1 t - 1 t − 1 個の 周辺分布の 比である( t = 1 t = 1 t = 1 では p ( x 1 ) p(x_1) p ( x 1 ) )。
証明. p ( x ≤ t ) ≥ p ( x 1 , … , x T ) > 0 p(x_{\leq t}) \geq p(x_1, \dots, x_T) > 0 p ( x ≤ t ) ≥ p ( x 1 , … , x T ) > 0 なので 比が 定義でき、積を とると 隣り合う 分母と 分子が 打ち消し合って p ( x ≤ T ) p(x_{\leq T}) p ( x ≤ T ) が 残る。 □ \square □
この 分解には 何の 仮定も 要らない。 自己回帰型の 言語モデル は、各条件付き分布を、前の トークンから ニューラルネットワーク(現在の 大規模言語モデルの 多くは トランスフォーマー。第5章 5.9 節)で 計算した 特徴量 h θ ( x < t ) h_\theta(x_{< t}) h θ ( x < t ) の ソフトマックス回帰 p θ ( x t ∣ x < t ) = softmax ( W h θ ( x < t ) ) x t p_\theta(x_t \mid x_{< t}) = \operatorname{softmax}(Wh_\theta(x_{< t}))_{x_t} p θ ( x t ∣ x < t ) = softmax ( W h θ ( x < t ) ) x t (第2章の 定義 2.19)で 表し、大規模言語モデルは これを 非常に 多くの パラメータと 大量の 文章で 学習した ものである。学習で 最小に する 平均の 交差エントロピー L ( θ ) = − 1 N ∑ t = 1 N log p θ ( x t ∣ x < t ) L(\theta) = -\frac{1}{N}\sum_{t=1}^{N}\log p_\theta(x_t \mid x_{< t}) L ( θ ) = − N 1 ∑ t = 1 N log p θ ( x t ∣ x < t ) (N N N は トークン数)は、命題 7.25 より 文章全体の 負の 対数尤度の 1 / N 1/N 1/ N なので、これは 最尤推定である(定理 7.7・命題 7.8)。
定義 7.26 (パープレキシティ, perplexity)評価用の 文章の N N N 個の トークンに ついて、 PPL = exp ( − 1 N ∑ t = 1 N log p θ ( x t ∣ x < t ) ) \operatorname{PPL} = \exp\bigl(-\frac{1}{N}\sum_{t=1}^{N}\log p_\theta(x_t \mid x_{< t})\bigr) PPL = exp ( − N 1 ∑ t = 1 N log p θ ( x t ∣ x < t ) ) を パープレキシティと いう。
パープレキシティは、正解の トークンに 与えた 確率の 逆数の 幾何平均で、 1 1 1 以上である。すべての トークンに 確率 1 / ∣ V ∣ 1/\lvert V \rvert 1/ ∣ V ∣ を 与える モデルでは ちょうど ∣ V ∣ \lvert V \rvert ∣ V ∣ なので、「平均して 何個の 候補の 間で 迷っているか」と 読める。文章の 生成は x t ∼ p θ ( ⋅ ∣ x < t ) x_t \sim p_\theta(\cdot \mid x_{< t}) x t ∼ p θ ( ⋅ ∣ x < t ) を 順に 生成して行う(温度 τ > 0 \tau > 0 τ > 0 を 使って softmax ( z / τ ) \operatorname{softmax}(z/\tau) softmax ( z / τ ) から 生成すると、 τ \tau τ が 小さい ほど 確率の 高いトークンに 集中する)。対話に 使う モデルでは、この 最尤推定(事前学習)の 後に、対話の 例に ついての 同じ 交差エントロピーの 最小化や、人の 評価を 使う 別の 目的関数などで 追加の 学習を するのが 一般的である。
注意
最尤推定が 求めるのは 学習データの 文章の 分布を まねる ことであり、生成された 文章の 内容が 事実と して 正しい ことは、目的関数に 直接は 入っていない。パープレキシティが 低いことも、内容の 正しさを 意味しない。
ヒント
実務では
パープレキシティは トークンあたりの 量なので、トークンへの 分け方が 違う モデルどうしでは 比べられない(細かく 分ける ほど 1 トークンあたりの 予測は 易しくなる)。比べるなら、同じ 評価用の 文章全体の 対数尤度を 文字数や バイト数で 割った 量を 使う(問題 7.7)。評価用の 文章が 学習データに 含まれていると 評価は 楽観的に なる( 第1章 の データ漏洩)。長い 文章の 確率は すぐに 浮動小数点数の 範囲を 下回るので、計算は つねに 対数で 行う。
まとめ
H ( p , q ) = H ( p ) + KL ( p ∥ q ) H(p, q) = H(p) + \operatorname{KL}(p \Vert q) H ( p , q ) = H ( p ) + KL ( p ∥ q ) であり、ギブスの 不等式 KL ( p ∥ q ) ≥ 0 \operatorname{KL}(p \Vert q) \geq 0 KL ( p ∥ q ) ≥ 0 (等号は p = q p = q p = q )は log u ≤ u − 1 \log u \leq u - 1 log u ≤ u − 1 から 従う。相互情報量は、独立な 場合の 分布からの KL ダイバージェンスである。
最尤推定は 経験分布との KL ダイバージェンス(交差エントロピー)の 最小化であり、交差エントロピー損失の ベイズ最適な 予測は 真の 条件付き分布である。
混合ガウスモデルの EM アルゴリズムは 負担率の 計算(E)と 重みつき最尤推定(M)の 繰り返しで、対数尤度は 単調に 増加する。しかし 最大でない 停留点で 止まる ことが あり、尤度は 非有界で、結果は 初期値に 依存する。
log p ( x ) = ELBO ( q ) + KL ( q ∥ p ( ⋅ ∣ x ) ) \log p(x) = \operatorname{ELBO}(q) + \operatorname{KL}(q \Vert p(\cdot \mid x)) log p ( x ) = ELBO ( q ) + KL ( q ∥ p ( ⋅ ∣ x )) 。平均場近似の 更新は q j ∝ exp ( E − j [ log p ( x , Z ) ] ) q_j \propto \exp(E_{-j}[\log p(x, Z)]) q j ∝ exp ( E − j [ log p ( x , Z )]) で、広がりを 過小評価しやすい。EM は q q q に 制約を 置かない ELBO の 座標上昇である。
隠れマルコフモデルの 尤度は 前向きアルゴリズムで O ( T K 2 ) O(TK^2) O ( T K 2 ) で 計算でき(素朴な 和は O ( T K T ) O(TK^T) O ( T K T ) )、状態の 列の 推定は ビタビ・アルゴリズム、パラメータの 推定は バウム–ウェルチ・アルゴリズム(EM)で 行う。
変分オートエンコーダは ELBO と 再パラメータ化、拡散モデルは 閉じた 形の 前向き過程と 雑音の 予測、大規模言語モデルは 自己回帰分解と 交差エントロピー(最尤推定)で 学習する。パープレキシティは 平均交差エントロピーの 指数である。
演習問題
問題 7.1 ★ (1) p = ( 1 / 2 , 1 / 4 , 1 / 4 ) p = (1/2, 1/4, 1/4) p = ( 1/2 , 1/4 , 1/4 ) と 一様分布 u = ( 1 / 3 , 1 / 3 , 1 / 3 ) u = (1/3, 1/3, 1/3) u = ( 1/3 , 1/3 , 1/3 ) に ついて、 H ( p ) H(p) H ( p ) 、KL ( p ∥ u ) \operatorname{KL}(p \Vert u) KL ( p ∥ u ) 、KL ( u ∥ p ) \operatorname{KL}(u \Vert p) KL ( u ∥ p ) を 求めよ。(2) 表の 確率が 0.1 , 0.5 , 0.9 0.1, 0.5, 0.9 0.1 , 0.5 , 0.9 の ベルヌーイ分布を p 1 , p 2 , p 3 p_1, p_2, p_3 p 1 , p 2 , p 3 と する。 KL ( p 1 ∥ p 3 ) > KL ( p 1 ∥ p 2 ) + KL ( p 2 ∥ p 3 ) \operatorname{KL}(p_1 \Vert p_3) > \operatorname{KL}(p_1 \Vert p_2) + \operatorname{KL}(p_2 \Vert p_3) KL ( p 1 ∥ p 3 ) > KL ( p 1 ∥ p 2 ) + KL ( p 2 ∥ p 3 ) を 確かめよ。
解答
(1) H ( p ) = 1 2 log 2 + 2 ⋅ 1 4 log 4 = 3 2 log 2 ≈ 1.0397 H(p) = \frac{1}{2}\log 2 + 2 \cdot \frac{1}{4}\log 4 = \frac{3}{2}\log 2 \approx 1.0397 H ( p ) = 2 1 log 2 + 2 ⋅ 4 1 log 4 = 2 3 log 2 ≈ 1.0397 、KL ( p ∥ u ) = log 3 − H ( p ) ≈ 0.0589 \operatorname{KL}(p \Vert u) = \log 3 - H(p) \approx 0.0589 KL ( p ∥ u ) = log 3 − H ( p ) ≈ 0.0589 (系 7.3 の 証明)、 KL ( u ∥ p ) = 1 3 ( log 2 3 + 2 log 4 3 ) = 1 3 log 32 27 ≈ 0.0566 \operatorname{KL}(u \Vert p) = \frac{1}{3}(\log\frac{2}{3} + 2\log\frac{4}{3}) = \frac{1}{3}\log\frac{32}{27} \approx 0.0566 KL ( u ∥ p ) = 3 1 ( log 3 2 + 2 log 3 4 ) = 3 1 log 27 32 ≈ 0.0566 。
(2) KL ( p 1 ∥ p 3 ) = 0.1 log 0.1 0.9 + 0.9 log 0.9 0.1 = 0.8 log 9 ≈ 1.7578 \operatorname{KL}(p_1 \Vert p_3) = 0.1\log\frac{0.1}{0.9} + 0.9\log\frac{0.9}{0.1} = 0.8\log 9 \approx 1.7578 KL ( p 1 ∥ p 3 ) = 0.1 log 0.9 0.1 + 0.9 log 0.1 0.9 = 0.8 log 9 ≈ 1.7578 。一方 KL ( p 1 ∥ p 2 ) = 0.1 log 0.2 + 0.9 log 1.8 ≈ 0.3681 \operatorname{KL}(p_1 \Vert p_2) = 0.1\log 0.2 + 0.9\log 1.8 \approx 0.3681 KL ( p 1 ∥ p 2 ) = 0.1 log 0.2 + 0.9 log 1.8 ≈ 0.3681 、KL ( p 2 ∥ p 3 ) = log 5 3 ≈ 0.5108 \operatorname{KL}(p_2 \Vert p_3) = \log\frac{5}{3} \approx 0.5108 KL ( p 2 ∥ p 3 ) = log 3 5 ≈ 0.5108 (例 7.4)で、和は 約 0.8789 0.8789 0.8789 に すぎない。
問題 7.2 ★ メールに 特定の 単語が 含まれるかを X X X (含めば 1 1 1 )、迷惑メールかを Y Y Y (迷惑メールなら 1 1 1 )とし、P ( X = 1 , Y = 1 ) = P ( X = 0 , Y = 0 ) = 0.4 P(X = 1, Y = 1) = P(X = 0, Y = 0) = 0.4 P ( X = 1 , Y = 1 ) = P ( X = 0 , Y = 0 ) = 0.4 、P ( X = 1 , Y = 0 ) = P ( X = 0 , Y = 1 ) = 0.1 P(X = 1, Y = 0) = P(X = 0, Y = 1) = 0.1 P ( X = 1 , Y = 0 ) = P ( X = 0 , Y = 1 ) = 0.1 と する。 I ( X ; Y ) I(X; Y) I ( X ; Y ) 、H ( X ) H(X) H ( X ) 、H ( X ∣ Y ) H(X \mid Y) H ( X ∣ Y ) を 求め、命題 7.6 の 2 を 確かめよ。
解答
周辺分布は どちらも ( 1 / 2 , 1 / 2 ) (1/2, 1/2) ( 1/2 , 1/2 ) なので H ( X ) = log 2 ≈ 0.6931 H(X) = \log 2 \approx 0.6931 H ( X ) = log 2 ≈ 0.6931 、I ( X ; Y ) = 0.8 log 0.4 0.25 + 0.2 log 0.1 0.25 = 0.8 log 1.6 + 0.2 log 0.4 ≈ 0.1927 I(X; Y) = 0.8\log\frac{0.4}{0.25} + 0.2\log\frac{0.1}{0.25} = 0.8\log 1.6 + 0.2\log 0.4 \approx 0.1927 I ( X ; Y ) = 0.8 log 0.25 0.4 + 0.2 log 0.25 0.1 = 0.8 log 1.6 + 0.2 log 0.4 ≈ 0.1927 。Y Y Y の どちらの 値のもとでも X X X の 条件付き分布は ( 0.8 , 0.2 ) (0.8, 0.2) ( 0.8 , 0.2 ) なので、H ( X ∣ Y ) = − 0.8 log 0.8 − 0.2 log 0.2 ≈ 0.5004 H(X \mid Y) = -0.8\log 0.8 - 0.2\log 0.2 \approx 0.5004 H ( X ∣ Y ) = − 0.8 log 0.8 − 0.2 log 0.2 ≈ 0.5004 で、H ( X ) − H ( X ∣ Y ) ≈ 0.1927 = I ( X ; Y ) H(X) - H(X \mid Y) \approx 0.1927 = I(X; Y) H ( X ) − H ( X ∣ Y ) ≈ 0.1927 = I ( X ; Y ) である。
問題 7.3 ★ ★ p p p を 平均 m m m 、分散 v > 0 v > 0 v > 0 の R \mathbb{R} R 上の 密度で、微分エントロピー h ( p ) = − ∫ p log p h(p) = -\int p\log p h ( p ) = − ∫ p log p が 有限な ものとする。 N ( μ , s 2 ) N(\mu, s^2) N ( μ , s 2 ) の うち KL ( p ∥ N ( μ , s 2 ) ) \operatorname{KL}(p \Vert N(\mu, s^2)) KL ( p ∥ N ( μ , s 2 )) を 最小に するのは μ = m \mu = m μ = m 、s 2 = v s^2 = v s 2 = v である ことを 示し、例 7.9 の N ( 0 , 10 ) N(0, 10) N ( 0 , 10 ) を 確かめよ。
解答
− log φ ( x ; μ , s 2 ) = 1 2 log ( 2 π s 2 ) + ( x − μ ) 2 2 s 2 -\log\varphi(x; \mu, s^2) = \frac{1}{2}\log(2\pi s^2) + \frac{(x - \mu)^2}{2s^2} − log φ ( x ; μ , s 2 ) = 2 1 log ( 2 π s 2 ) + 2 s 2 ( x − μ ) 2 と E p [ ( X − μ ) 2 ] = v + ( m − μ ) 2 E_p[(X - \mu)^2] = v + (m - \mu)^2 E p [( X − μ ) 2 ] = v + ( m − μ ) 2 より
KL ( p ∥ N ( μ , s 2 ) ) = − h ( p ) + 1 2 log ( 2 π s 2 ) + v + ( m − μ ) 2 2 s 2 \operatorname{KL}(p \Vert N(\mu, s^2)) = -h(p) + \frac{1}{2}\log(2\pi s^2) + \frac{v + (m - \mu)^2}{2s^2} KL ( p ∥ N ( μ , s 2 )) = − h ( p ) + 2 1 log ( 2 π s 2 ) + 2 s 2 v + ( m − μ ) 2
で、これは μ = m \mu = m μ = m でだけ 最小に なる。その とき u = s 2 u = s^2 u = s 2 の 関数 1 2 log u + v 2 u \frac{1}{2}\log u + \frac{v}{2u} 2 1 log u + 2 u v の 導関数 u − v 2 u 2 \frac{u - v}{2u^2} 2 u 2 u − v は u = v u = v u = v の 前後で 負から 正に 変わる。例 7.9 の 混合分布の 平均は 0 0 0 、分散は E [ X 2 ] = 1 + 9 = 10 E[X^2] = 1 + 9 = 10 E [ X 2 ] = 1 + 9 = 10 である。
問題 7.4 ★ ★ ある 分析者が、顧客の 購入額の 対数に 混合ガウスモデルを 当てはめ、成分の 数 K = 1 , … , 8 K = 1, \dots, 8 K = 1 , … , 8 の それぞれで EM アルゴリズムを 1 回ずつ 実行した。訓練データの 対数尤度は K K K とともに おおむね増え、 K = 8 K = 8 K = 8 で ほかより 格段に 大きくなった。 K = 8 K = 8 K = 8 の 解では、ある 成分の 分散が 10 − 9 10^{-9} 1 0 − 9 で、その 重みは 1 / n 1/n 1/ n に 近かった。分析者は「顧客は 8 つの 層に 分かれる」と 結論した。この 分析の 問題点を 挙げ、どう すべきかを 述べよ。
解答
(1) 訓練データの 対数尤度は、成分を 複製すれば K K K に ついて 下がらないので、 K K K の 選択に 使えない。(2) 分散が 10 − 9 10^{-9} 1 0 − 9 で 重みが 1 / n 1/n 1/ n に 近い 成分は ほぼ 1 点に 縮退しており、格段に 大きい 対数尤度は 命題 7.15 の 特異点に よる 見かけの ものである(同じ 購入額が 繰り返し現れると 起こりやすい)。(3) 各 K K K で 1 回しか 実行しておらず、局所解や 初期値への 依存を 確かめていない。(4) 成分は 顧客の 層とは 限らない。山が 1 つでも、歪んだ 分 布や 裾の 重い 分布は いくつもの 正規分布の 和で 近似される。分散に 下限を 設け、初期値を 変えて 何回も 実行し、 K K K は 検証データでの 対数尤度や BIC で 選び、解の 安定性と、成分が 業務上意味の ある 違いに 対応するかを 確かめるべきである。
問題 7.5 ★ ★ 隠れマルコフモデルで β T ( k ) = 1 \beta_T(k) = 1 β T ( k ) = 1 、t < T t < T t < T に ついて β t ( k ) = ∑ z t + 1 , … , z T ∏ s = t + 1 T a z s − 1 z s b z s ( x s ) \beta_t(k) = \sum_{z_{t+1}, \dots, z_T}\prod_{s=t+1}^{T}a_{z_{s-1}z_s}b_{z_s}(x_s) β t ( k ) = ∑ z t + 1 , … , z T ∏ s = t + 1 T a z s − 1 z s b z s ( x s ) (z t = k z_t = k z t = k と する)と 定める。(1) β t ( j ) = ∑ k a j k b k ( x t + 1 ) β t + 1 ( k ) \beta_t(j) = \sum_ka_{jk}b_k(x_{t+1})\beta_{t+1}(k) β t ( j ) = ∑ k a j k b k ( x t + 1 ) β t + 1 ( k ) と p ( x , Z t = k ) = α t ( k ) β t ( k ) p(x, Z_t = k) = \alpha_t(k)\beta_t(k) p ( x , Z t = k ) = α t ( k ) β t ( k ) を 示せ。(2) 例 7.22 で β 1 , β 2 \beta_1, \beta_2 β 1 , β 2 と P ( Z 2 = 2 ∣ x ) P(Z_2 = 2 \mid x) P ( Z 2 = 2 ∣ x ) を 求め、時刻 2 の フィルタリングの 値と 比べよ。
解答
(1) β t ( j ) \beta_t(j) β t ( j ) の 定義の 和で z t + 1 = k z_{t+1} = k z t + 1 = k を 固定すると 因子 a j k b k ( x t + 1 ) a_{jk}b_k(x_{t+1}) a j k b k ( x t + 1 ) が 外に 出て、残りの 和が β t + 1 ( k ) \beta_{t+1}(k) β t + 1 ( k ) に なる。 z t = k z_t = k z t = k の とき p ( x , z ) p(x, z) p ( x , z ) は s ≤ t s \leq t s ≤ t の 因子の 積((2) の 右辺)と s > t s > t s > t の 因子の 積の 積であり、 z 1 : t − 1 z_{1:t-1} z 1 : t − 1 に ついて 和を とると 前者は α t ( k ) \alpha_t(k) α t ( k ) 、z t + 1 : T z_{t+1:T} z t + 1 : T に ついて 和を とると 後者は β t ( k ) \beta_t(k) β t ( k ) に なる。
(2) β 2 ( 1 ) = 0.9 × 0.1 + 0.1 × 0.6 = 0.15 \beta_2(1) = 0.9 \times 0.1 + 0.1 \times 0.6 = 0.15 β 2 ( 1 ) = 0.9 × 0.1 + 0.1 × 0.6 = 0.15 、β 2 ( 2 ) = 0.3 × 0.1 + 0.7 × 0.6 = 0.45 \beta_2(2) = 0.3 \times 0.1 + 0.7 \times 0.6 = 0.45 β 2 ( 2 ) = 0.3 × 0.1 + 0.7 × 0.6 = 0.45 、β 1 ( 1 ) = 0.9 × 0.1 × 0.15 + 0.1 × 0.6 × 0.45 = 0.0405 \beta_1(1) = 0.9 \times 0.1 \times 0.15 + 0.1 \times 0.6 \times 0.45 = 0.0405 β 1 ( 1 ) = 0.9 × 0.1 × 0.15 + 0.1 × 0.6 × 0.45 = 0.0405 、β 1 ( 2 ) = 0.3 × 0.1 × 0.15 + 0.7 × 0.6 × 0.45 = 0.1935 \beta_1(2) = 0.3 \times 0.1 \times 0.15 + 0.7 \times 0.6 \times 0.45 = 0.1935 β 1 ( 2 ) = 0.3 × 0.1 × 0.15 + 0.7 × 0.6 × 0.45 = 0.1935 (検算:0.81 × 0.0405 + 0.04 × 0.1935 = 0.040545 = p ( x ) 0.81 \times 0.0405 + 0.04 \times 0.1935 = 0.040545 = p(x) 0.81 × 0.0405 + 0.04 × 0.1935 = 0.040545 = p ( x ) )。P ( Z 2 = 2 ∣ x ) = 0.0654 × 0.45 / 0.040545 ≈ 0.726 P(Z_2 = 2 \mid x) = 0.0654 \times 0.45/0.040545 \approx 0.726 P ( Z 2 = 2 ∣ x ) = 0.0654 × 0.45/0.040545 ≈ 0.726 で、時刻 2 までの 観測に よる フィルタリングの 値 0.469 0.469 0.469 より 大きい。時刻 3 の アラームが、時刻 2 の 異常の 証拠を 強めたのである。
問題 7.6 ★ ★ ★ X 0 X_0 X 0 を 密度 p 0 p_0 p 0 を もつ R d \mathbb{R}^d R d の 確率ベクトル、 ε ∼ N d ( 0 , I ) \varepsilon \sim N_d(0, I) ε ∼ N d ( 0 , I ) を X 0 X_0 X 0 と 独立とし、 0 < α ˉ < 1 0 < \bar{\alpha} < 1 0 < α ˉ < 1 に ついて X = α ˉ X 0 + 1 − α ˉ ε X = \sqrt{\bar{\alpha}}X_0 + \sqrt{1 - \bar{\alpha}}\varepsilon X = α ˉ X 0 + 1 − α ˉ ε の 密度を q q q と する。微分と 積分の 順序交換を 認めて、 E [ ε ∣ X = x ] = − 1 − α ˉ ∇ log q ( x ) E[\varepsilon \mid X = x] = -\sqrt{1 - \bar{\alpha}}\nabla\log q(x) E [ ε ∣ X = x ] = − 1 − α ˉ ∇ log q ( x ) (ツイーディーの 公式)を 示せ。
解答
X 0 = x 0 X_0 = x_0 X 0 = x 0 のもとでの X X X の 密度 k ( x ∣ x 0 ) = φ ( x ; α ˉ x 0 , ( 1 − α ˉ ) I ) k(x \mid x_0) = \varphi(x; \sqrt{\bar{\alpha}}x_0, (1 - \bar{\alpha})I) k ( x ∣ x 0 ) = φ ( x ; α ˉ x 0 , ( 1 − α ˉ ) I ) に ついて q ( x ) = ∫ k ( x ∣ x 0 ) p 0 ( x 0 ) d x 0 q(x) = \int k(x \mid x_0)p_0(x_0)\ dx_0 q ( x ) = ∫ k ( x ∣ x 0 ) p 0 ( x 0 ) d x 0 、∇ x k ( x ∣ x 0 ) = − x − α ˉ x 0 1 − α ˉ k ( x ∣ x 0 ) \nabla_xk(x \mid x_0) = -\frac{x - \sqrt{\bar{\alpha}}x_0}{1 - \bar{\alpha}}k(x \mid x_0) ∇ x k ( x ∣ x 0 ) = − 1 − α ˉ x − α ˉ x 0 k ( x ∣ x 0 ) である。k ( x ∣ x 0 ) p 0 ( x 0 ) / q ( x ) k(x \mid x_0)p_0(x_0)/q(x) k ( x ∣ x 0 ) p 0 ( x 0 ) / q ( x ) は X = x X = x X = x のもとでの X 0 X_0 X 0 の 条件付き密度なので
∇ log q ( x ) = ∇ q ( x ) q ( x ) = − 1 1 − α ˉ ∫ ( x − α ˉ x 0 ) k ( x ∣ x 0 ) p 0 ( x 0 ) q ( x ) d x 0 = − 1 1 − α ˉ E [ X − α ˉ X 0 ∣ X = x ] \nabla\log q(x) = \frac{\nabla q(x)}{q(x)} = -\frac{1}{1 - \bar{\alpha}}\int(x - \sqrt{\bar{\alpha}}\,x_0)\frac{k(x \mid x_0)p_0(x_0)}{q(x)}\ dx_0 = -\frac{1}{1 - \bar{\alpha}}E\bigl[X - \sqrt{\bar{\alpha}}\,X_0 \bigm| X = x\bigr] ∇ log q ( x ) = q ( x ) ∇ q ( x ) = − 1 − α ˉ 1 ∫ ( x − α ˉ x 0 ) q ( x ) k ( x ∣ x 0 ) p 0 ( x 0 ) d x 0 = − 1 − α ˉ 1 E [ X − α ˉ X 0 X = x ]
で、X − α ˉ X 0 = 1 − α ˉ ε X - \sqrt{\bar{\alpha}}X_0 = \sqrt{1 - \bar{\alpha}}\varepsilon X − α ˉ X 0 = 1 − α ˉ ε を 代入すればよい。雑音を 予測する ネットワークは、雑音を 加えた データの 分布の スコアを 推定しているのである。
問題 7.7 ★ ★ (1) ある 言語モデルが 評価用の 文章の 4 個の トークンに 与えた 確率が 順に 1 / 2 , 1 / 4 , 1 / 8 , 1 / 2 1/2, 1/4, 1/8, 1/2 1/2 , 1/4 , 1/8 , 1/2 だった。平均交差エントロピーと パープレキシティを 求めよ。(2) 同じ 評価用の 文章で、語彙が 32000 の モデル A の パープレキシティは 20、1 文字を 1 トークンと する モデル B の パープレキシティは 8 だった。A の 1 トークンは 平均 1.5 文字に あたる。「B の ほうが 文章を よく 予測している」と いう 結論は 正しいか。
解答
(1) 平均交差エントロピーは 1 4 ( log 2 + log 4 + log 8 + log 2 ) = 7 4 log 2 ≈ 1.213 \frac{1}{4}(\log 2 + \log 4 + \log 8 + \log 2) = \frac{7}{4}\log 2 \approx 1.213 4 1 ( log 2 + log 4 + log 8 + log 2 ) = 4 7 log 2 ≈ 1.213 、パープレキシティは 2 7 / 4 ≈ 3.364 2^{7/4} \approx 3.364 2 7/4 ≈ 3.364 。
(2) 正しくない。文章の 文字数を C C C と すると、文章全体の 対数尤度は A では − C 1.5 log 20 -\frac{C}{1.5}\log 20 − 1.5 C log 20 、B では − C log 8 -C\log 8 − C log 8 で、1 文字あたり A は 約 − 1.997 -1.997 − 1.997 、B は 約 − 2.079 -2.079 − 2.079 である。A の ほうが 文章全体に 高い 確率を 与えており(1 文字あたりの パープレキシティは 20 2 / 3 ≈ 7.37 < 8 20^{2/3} \approx 7.37 < 8 2 0 2/3 ≈ 7.37 < 8 )、よく 予測している。トークンへの 分け方が 違う モデルを、トークンあたりの パープレキシティで 比べてはいけない。