交差エントロピー | 情報理論
交差エントロピー $H(p,q)$ は、確率分布 $q$ から発生する情報量(符号長)$I(x) = -\log_2 q(x)$ を、真の確率 $p(x)$ で重み付けして平均をとったものです。つまり $q$ の情報量(符号長)が $p$ に従って発生したときに必要な平均ビット数を表しています。
\[\begin{aligned} H(p,q) = - \sum_{x} p(x) \log_2 q(x) \quad \text{離散確率変数} \\ H(p,q) = - \int\limits_x p(x) \log_2 q(x) dx \quad \text{連続確率変数} \\ \end{aligned}\]期待値を使って以下のように定義できます。
\[\mathbb{E}_{x \sim p}[-\log_2 q(x)]\]$p=q$ なら、割り当てた符号長がその事象の真の分布と合っているので、平均符号長は最小の $H(p)$ になります($H(p,q)=H(p)$)。
$p \neq q$ だと、符号長の割り当てが実際の起こりやすさと乖離しているため平均符号長は $H(p)$ より大きくなります。 その差が KLダイバージェンス($D_{\text{KL}}(p\parallel q)$)です。
交差エントロピー $H(p,q)$ は KLダイバージェンス($D_{\text{KL}}(p\parallel q)$)を用いて以下のように表すことができます。
\[\begin{aligned} H(p, q) = H(p) + D_{\text{KL}}(p \parallel q) \\ D_{\text{KL}}(p \parallel q) = H(p,q) - H(p) \geq 0 \\ p=q \Rightarrow D_{\text{KL}}(p\parallel q) = H(p,q)=H(p) = 0 \end{aligned}\]交差エントロピー誤差(損失関数・誤差関数)
交差エントロピーは、2 つの確率分布の乖離を評価できるため、Deep Learning の分類問題の損失(誤差)関数(Loss Function)として使用されます(誤差関数を強調する意味で交差エントロピー誤差(Cross-Entropy Error / Cross-Entropy Loss)と呼ぶことが多いようです)。
- $p(x)$(現実の分布):正解データ(正解ラベル)の分布
- $q(x)$(モデルの分布):モデルが出力した予測(Softmax 関数を通した確率)の分布
[!NOTE] Deep Learning の実装(PyTorchなど)では、情報理論で一般的な底が $2$ の対数($\log_2$)ではなく、自然対数($\log_e$ または $\ln$)が使われます。以降の計算例では自然対数を使用します。
コンテキスト長 $C$ 内の「ある 1 つの位置(トークン位置)$t$」に注目したときの、Deep Learning のモデルの分布と正解データの関係性は以下の通りです。
1. モデルの出力(予測確率分布: $q_t$)
各位置 $t$ において、モデルは次に続くトークンの予測として、語彙数(Vocabulary Size: $V$)と同じ次元数を持つ確率分布を出力します。
\[q_t^i = (q_t^1,\ q_t^2,\ \dots,\ q_t^V), \quad \sum_{i=1}^V q_t^i = 1\]2. 正解データ(one-hotベクトル: $p_t$)
正解データ(次に実際に来た単語)を、語彙数 $V$ の次元を持つ One-Hotベクトル(正解の単語のインデックスだけが 1 で、他はすべて 0 のベクトル)として表現します。
3. 交差エントロピー誤差($L_t$)の計算
\[L_t = - \sum_{i=1}^{V} p_t(i) \log q_t(i)\]正解の one-hotベクトル $p_t$ は正解ラベル($k_t$ 番目)以外すべて $0$ になるため、正解トークンの項だけが残ります。
\[L_t = - \log q_t(k_t)\]4. コンテキスト(系列長 C )全体の損失
各 $t$ の交差エントロピー誤差を算出して平均をとります。
これは系列長 C 全体の尤度 $\prod_t q_t(k_t)$ の対数を取り符号を反転させたものと一致します。
ref. 系列長 C の尤度関数 / 対数尤度関数 と損失関数
5. 全体(バッチとコンテキスト)での損失の集約
LLM の学習は入力データのシェイプであるバッチ数 $N$ とコンテキスト長 $C$ の損失を一括で計算します。
- すべての位置での計算: バッチ内の各サンプル $n$($1$ から $N$)の、各トークン位置 $t$($1$ から $C$)のすべてにおいて、上記の一連の交差エントロピー損失 $L_{n, t}$ を計算します。
- 平均値の算出: すべての位置の損失の平均値をとり、全体の総損失(Total Loss)とします。
つまり総損失(Total Loss)を交差エントロピー誤差の平均値と定義します。
\[\text{Total Loss} = \frac{1}{N \times C} \sum_{n=1}^{N} \sum_{t=1}^{C} L_{n, t} = \frac{1}{N \times C} \sum_{n=1}^{N} \sum_{t=1}^{C} -\log q_t(k_t)\]この全体の平均損失(Total Loss)が小さくなるように、バックプロパゲーション(誤差逆伝播法)によって Transformer 内部のパラメータ(重み)が更新されていきます。
モンテカルロ積分によるクロスエントロピーの近似
クロスエントロピー $H(p,q) = - \int p(x) \log_2 q(x) dx$ を解析的に求めることが困難な場合、モンテカルロ積分を用いて近似します。
モンテカルロ積分では、真の分布 $p(x)$ による期待値を、真の分布から抽出されたデータ $D={x_{1},\ldots,x_{n}}$(ただし $x_i \sim p(x)$)の標本平均に置き換えます。
\[H(p,q) \approx - \frac{1}{n} \sum_{i=1}^n \log_2 q(x_i)\]データ数 $n$ で割ることで、サンプルの平均(期待値の近似値)を求めます。
対数の性質を利用すると、以下のように変形することもできます。
\[\begin{aligned} H(p,q) &\approx - \frac{1}{n} \sum_{i=1}^n \log_2 q(x_i) \\ &= - \frac{1}{n} \log_2 \prod_{i=1}^n q(x_i) \end{aligned}\]Python コード
\[L = -\sum_{k=0}^{n} t_k \log y_k\](本文中の記法とは異なり、この式のみ $0$ 始まり・上限 $n$ の慣習で書いています。)
def cross_entropy_error(y, t):
if y.ndim == 1:
t = t.reshape(1, t.size)
y = y.reshape(1, y.size)
# 正解データがone-hot-vectorの場合、正解ラベルのインデックスに変換
if t.size == y.size:
t = t.argmax(axis=1)
batch_size = y.shape[0]
# np.log(0) はマイナス無限大(-inf)になり、計算が止まってしまいます。
# それを防ぐために、ごく小さな値(1e-7)を足して、絶対に 0 にならないように安全策をとっています。
return -np.sum(np.log(y[np.arange(batch_size), t] + 1e-7)) / batch_size
出典:ゼロから作るDeep Learning ❷ ―自然言語処理編 斎藤康毅 (著)