ゼロから作るDeep Learning ❷ ―自然言語処理編 斎藤康毅 (著) に登場する TimeSoftmaxWithLoss レイヤの逆伝播に関するメモです。

記法

本ドキュメントは以下の記法で統一します。

記号 意味
$\cdot$ 行列積(ドット積)
$\odot$ 要素積(アダマール積、element-wise product)

前提と記号の定義

  • corpus size: 1000
  • V( vocab_size ):コーパスから重複を排除した単語 ID の数(語彙数)
    • ソースコードは 418
  • N:ミニバッチサイズ
    • ソースコードは 10
  • T:系列長(sequence length)
    • ソースコードは 5
  • E( wordvec_size ):単語の分散表現の次元数(要素数)
    • ソースコードは 100
  • H( hidden_size ): RNNの隠れ状態の次元数(要素数)
    • ソースコードは 100

※ ソースコードの E, H はともに 100 ですが異なる値を設定できます。

  • $\mathbf{s}$:Softmax 関数に渡される前のスコア $\mathbf{s} = (s_1, s_2, \dots, s_k, \dots, s_V)$ (Softmax 関数 に渡されるミニバッチ $(N, T, V)$ の $N=n, T=t$ のベクトル$(V,)$)
    • Softmax に渡されるミニバッチは (Time)Affine レイヤで t 時刻の隠れ状態 $h_t$ から算出した次にくる単語スコアのミニバッチでシェイプは $(N, T, V)$
  • $\mathbf{q}$:Softmax 後の確率分布 $\mathbf{q} =(q_1, q_2, \dots, q_k, \dots, q_V)$
    • $q_k = \dfrac{\exp(s_k)}{\sum_{i=1}^V \exp(s_i)}$
  • $\mathbf{y}$ = 正解データの one-hot ベクトル $\mathbf{y} = (y_1, y_2, \dots, y_k, \dots, y_V)$
    • $y_k$ : 正解クラスを $k$ とすると $y_k=1$、それ以外は $0$
  • $L = -\sum_{i=1}^V y_i \log q_i$ : 交差エントロピー誤差。$\mathbf{y}$ が one-hot なので実質 $-\log q_k$ と同じ (※ k が正解クラスの場合)

逆伝播 $\dfrac{\partial L}{\partial s} = q - y$

$\dfrac{\partial L}{\partial s}$(Softmax に入る前のスコアに対する勾配)を求めます。

導出(対数の性質を使う方法)

まず $\log q_k$ を展開します。

\[\log q_k = \log \frac{\exp(s_k)}{\sum_j \exp(s_j)} = s_k - \log\sum_j \exp(s_j)\]

これを $L$(交差エントロピー誤差)に代入すると

\[L = -\sum_k y_k \left( s_k - \log\sum_j \exp(s_j) \right) = -\sum_k y_k s_k \;+\; \left(\log\sum_j \exp(s_j)\right)\sum_k y_k\]

$y$ は one-hot なので $\sum_k y_k = 1$。よって

\[L = -\sum_k y_k s_k + \log\sum_j \exp(s_j)\]

この形は $s_i$ に対して微分しやすくなっています。

第1項の微分:

\[\frac{\partial}{\partial s_i}\left(-\sum_k y_k s_k\right) = -y_i\]

($s_k$ たちは独立変数なので、$k=i$ の項だけ残る)

第2項の微分:

\[\frac{\partial}{\partial s_i}\log\sum_j \exp(s_j) = \frac{\exp(s_i)}{\sum_j \exp(s_j)} = q_i\]

合わせると:

\[\frac{\partial L}{\partial s_i} = -y_i + q_i = q_i - y_i\]

これがまさに $q - y$ です。

直感的な意味

  • 正解クラス($y_i=1$)では勾配は $q_i - 1$。予測確率 $q_i$ が $1$ に近ければ勾配は $0$ に近づき、遠ければ大きな勾配(=大きな修正圧力)がかかる。
  • 不正解クラス($y_i=0$)では勾配は $q_i$ そのもの。誤って高い確率を割り当てているほど、その確率を下げる方向に大きく修正される。

Softmax と交差エントロピーを組み合わせると、途中の複雑な微分(Softmax の分数の微分など)がきれいに打ち消し合って、こんなにシンプルな $q-y$ という形になる、というのがこの式のありがたいところです。

例えば $\mathbf{q}$ が [0.23, 0.05, 0.72] で正解 $y$( one-hot )が [0, 0, 1] なら [0.23, 0.05, -0.28] が勾配になります。

補足(別解:Softmax のヤコビアンから)

もう一つのやり方として、Softmax 自体の微分 $\dfrac{\partial q_k}{\partial s_i} = q_i(\delta_{ik} - q_k)$($\delta_{ik}$ はクロネッカーのデルタ)を使い、$L=-\log q_t$ をチェーンルールで微分しても同じ結果 $q_i - y_i$ にたどり着きます。ただし対数トリックを使う上の方法の方が計算がシンプルです。

なお実装上(TimeSoftmaxWithLoss)では、この $q-y$ をバッチサイズ(と有効なタイムステップ数)で割って平均化してから逆伝播させている点も、コードを読む際は注意してください。

実装との対応

ゼロから作るDeep Learning ❷ ―自然言語処理編 斎藤康毅 (著) の実装は以下のとおりです。

dx = mask.sum()
# ys.shape = (N*T, V)
# ys は次にくる単語のスコア( logits )
ts, ys, mask, (N, T, V) = self.cache

# ts.shape = (N*T, ) 教師データの単語 ID
# ys.shape = (N*T, V)
dx = ys

# W := W − η * (∂W / ∂L)
# 正解    <= 0 なので重みは増加
# 不正解  >  0 なので重みは減少
dx[np.arange(N * T), ts] -= 1
dx *= dout
dx /= mask.sum() # /= は「除算代入演算子」で、dx = dx / mask.sum() と同じ意味
dx *= mask[:, np.newaxis]  # ignore_labelに該当するデータは勾配を0にする

dx = dx.reshape((N, T, V))

return dx

![NOTE]
実際のコードは汎用性を持たせるために ignore-label の処理をするが上記コードでは 簡単にするために無視する( loss = mask.sum() にしている) #ls = np.log(ys[np.arange(N * T), ts]) ls *= mask # ignore_labelに該当するデータは損失を0にする loss = -np.sum(ls) loss /= mask.sum()

以下で $p - y$ を算出します。

dx[np.arange(N*T), ts] -= 1
  • 正解:p_i -> p_i - 1
  • 正解以外: p_i

その後、dx /= mask.sum() によって平均損失に対応する勾配へ変換します。

dx = np.array([
   [0.23, 0.05, 0.72],
   [0.75, 0.20, 0.05],
   [0.40, 0.50, 0.10]]);

# dx[[1,2]] 行に関して抽出
[[0.75, 0.20 , 0.05],
[0.40 , 0.50 , 0.10 ]]

# 対象の要要を一括で抽出
dx[[1,2],[0,1]]
[0.75, 050] 

# numpy.arange 関数
# numpy.arange(start, stop, step) は等間隔の数値配列を生成する関数です。

np.arange(5)          # [0, 1, 2, 3, 4]
np.arange(1, 5)       # [1, 2, 3, 4]
np.arange(0, 10, 2)   # [0, 2, 4, 6, 8]
np.arange(0, 1, 0.25) # [0, 0.25, 0.5, 0.75] 

#stopは含まれません(半開区間)。
# 浮動小数点を使うと誤差が出ることがあるため、その場合はnumpy.linspaceの使用が推奨されます。

例えば予測確率 $p$ が [0.23, 0.05, 0.72] で正解が 2 なら $y$ は [0, 0, 1] なので $p - y$ は [0.23, 0.05, -0.28] になり正解クラスだけ負になります。

SGD の更新式:

\[W := W - \eta \frac{\partial L}{\partial W}\]

正解クラス

勾配が負 $-0.28$ なので $W - η(-0.28) = W + η0.28$ となり、そのクラスのスコア( logit )が上がる方向に更新されます。

結果として次回の予測では正解クラスの確率が高くなります。

不正解クラス

例えばクラス0の勾配は $+0.23$ です。
更新すると $W - η(0.23)$ となるため、そのクラスのスコアは下がる方向へ更新されます。

結果として不正解クラスの確率は低くなります。

  • 正解クラス $\rightarrow$ 勾配が負 $\rightarrow$ スコアを上げる方向に更新
  • 不正解クラス $\rightarrow$ 勾配が正 $\rightarrow$ スコアを下げる方向に更新

そのため $p - y$ は「正解クラスの確率を上げ、不正解クラスの確率を下げるための誤差信号」として解釈できます。

shape の流れ

処理 shape
TimeAffine出力 $(N, T , V)$
reshape後 $(N * T, V)$
Softmax出力 $(N * T, V)$
backwardの勾配 $(N * T, V)$
reshapeして戻す $(N, T, V)$

上記より TimeSoftmaxWithLoss が返す勾配はそのまま TimeAffine.backward() に渡されます。