LLM の条件付き確率分布と尤度 | 確率・統計
本稿では、条件付き確率分布の考え方を自己回帰言語モデル(LLM)に適用し、連鎖律による同時分布の分解から、学習時の誤差関数(交差エントロピー / 負の対数尤度)までを一貫した流れで整理します。
以降では、モデルが近似する分布を $q$、真の(データが従う)分布を $p$ として区別して表記します。
LLMにおける確率変数と実現値
- $X_1, X_2, \dots, X_T$:各位置のトークンを表す確率変数
- $x_1, x_2, \dots, x_T$:その実現値(具体的な単語・トークン)
- $x_{<t} := (x_1, \dots, x_{t-1})$:$t$ 番目より前の実現値をまとめた記法
条件付き確率 $q(x_t\vert{}x_{<t})$
条件付き確率 $q(x_t\vert{}x_{
\[q(x_t\vert{}x_{<t}) := q(X_t = x_t\vert{}X_1=x_1,\dots,X_{t-1}=x_{t-1}) = \frac{q(x_1,\dots,x_t)}{q(x_1,\dots,x_{t-1})}\]条件 $x_{<t}$ は $t-1$ 個の確率変数それぞれに関する事象の連言を圧縮して書いたものであり、これらの確率変数の存在を前提にしないと厳密には意味が取れません。
条件付き確率分布としての意味
$x_{<t}$ を固定して、$x_t$ が語彙 $V$ の中のどの値を取るかについて確率を割り当てたもの全体
\[\{\, q(x_t = v\vert{}x_{<t}) \;:\; v \in V \,\}\]が「トークン $X_t$ の、$x_{
系列全体への拡張(連鎖律)
条件付き確率の定義(乗法定理)
\[P(A,B) = P(A)\,P(B\mid A)\]を $T$ 個の確率変数 $x_1,\ldots,x_T$ に逐次的に適用すると、系列全体の同時分布は
\[q(x_1,\dots,x_T) = q(x_1)\, q(x_2\mid x_1)\, q(x_3\mid x_1,x_2) \cdots q(x_T\vert{}x_{<T}) = \prod_{t=1}^{T} q(x_t\vert{}x_{<t})\]と分解できます。これは確率の公理から導かれる恒等式(連鎖律, chain rule of probability)であり、i.i.d.(独立同分布)の仮定を一切必要としません。各因子 $q(x_t\mid x_{
- $t$ ごとに異なる分布であってよい(identically distributed である必要はない)
- 直前の文脈 $x_{<t}$ に依存してよい(independent である必要はない)
という点で、i.i.d. を仮定した単純な積 $\prod_t q(x_t)$ とは前提が異なります。自己回帰言語モデルは、この連鎖律をそのまま利用して、各時刻 $t$ で文脈 $x_{
LLM の学習(自己回帰的な尤度最大化)は、この各項 $q(x_t\mid x_{
LLMの交差エントロピー誤差とカテゴリカル分布の尤度関数
ミニバッチサイズ $N$、系列長 $C$、語彙数 $V$ とします。ここでは 1 バッチ(例 $N=1$)、すなわち系列長 $C$ の 1 系列に注目して、交差エントロピー誤差と負の対数尤度(NLL: Negative Log-Likelihood)の対応を確認します。バッチ次元 $N$ 方向の平均は、この 1 系列分の Total Loss をさらにサンプル間で平均したものに相当します。
各位置 $t$ の予測分布 $q_t$ は、$x_{
$x_{t,v} \in {0,1}$ は時刻 $t$ における正解トークンを表す one-hot ベクトルの $v$ 番目の成分を表します。
時刻 $t$ の正解ラベルが $k_t$ であるとき、$x_{t,v}$ は次式で定義されます。
このとき時刻 $t$ の尤度関数は 観測値が1つ(正解トークン) のカテゴリカル分布の尤度関数として次式で定義されます。
\[\prod_{v=1}^{V} q_t(v)^{x_{t,v}} = q_t(k_t)\]$v=k_t$ の項以外はすべて $q_t(v)^0=1$ となるため、実質的に正解トークンの確率のみが残ります。
また対数尤度は次式で定義されます。
系列長 C の尤度関数 / 対数尤度関数 と損失関数
系列長 $C$ 全体の尤度関数、対数尤度は、前節の連鎖律により(各時刻で異なる分布 $q_t$ であっても成立する形で)以下の式で定義されます。
\[\begin{aligned} \prod_{t=1}^{C} q_t(k_t) \quad 系列長 C の尤度関数 \\ \sum_{t=1}^C \log q_t \left(k_t \right) \quad 系列長 C の対数尤度関数 \end{aligned}\]LLM の系列長 C の損失(交差誤差エントロピーの平均)は 列長 $C$ の対数尤度の(トークン数で正規化した)平均を取り符号を反転させて以下の式で定義されます。
\[L = - \frac{1}{C} \sum_{t=1}^C \log q_t \left(k_t \right)\]各 $t$ における分布 $q_t$ はそれぞれ独自の分布、つまり位置ごとに異なる分布 $q_1, \ldots, q_t, \ldots, q_C$ からそれぞれ 1 回ずつサンプリングした結果を掛け合わせたものです。$q_t$ はモデルが文脈(それまでのトークン)を条件として出力する条件付き分布なので、$t$ が変われば分布のパラメータ自体が変わります。
厳密に書くと以下になります。
\[\prod_{t=1}^{C} q(x_t\vert{}x_{<t}) = \prod_{t=1}^{C} q_t(k_t)\]自己回帰言語モデルの交差エントロピー誤差は、系列全体の尤度 $\prod_t q_t(k_t)$ の対数を取り符号を反転させたものであり、連鎖律によって保証された厳密な同時分布の分解に基づいています。
これが、交差エントロピー誤差を用いた学習が言語モデルの尤度最大化と数学的に等価であることの根拠になっています。
まとめ
| 要素 | 対応するもの |
|---|---|
| 確率変数 | $X_1,\dots,X_T$(各位置のトークン) |
| 実現値 | $x_1,\dots,x_T$(実際の文) |
| 条件(複合事象) | $x_{<t}$(それまでの文脈) |
| 条件付き確率分布 | $q(x_t\mid x_{<t})$(語彙上の分布、モデルの出力) |
| 同時分布の分解 | 連鎖律 $\prod_t q(x_t\mid x_{ |
| 系列の尤度 | $\prod_{t=1}^{C} q_t(k_t)$ |
| 対数尤度(平均) | $\displaystyle\frac{1}{C}\sum_{t=1}^{C}\log q_t(k_t)$ |
| 学習時の誤差関数 | 交差エントロピー誤差 = 負の対数尤度(NLL) |