ミニプロジェクト:Self-Attention を自分で作る
Transformer の心臓部である、マスク付きの Self-Attention を、NumPy だけで作ります。
作るもの
入力は、3つの単語の埋め込みベクトル(それぞれ2個の数)を並べた表 (3行 × 2列)です。
- 、、 で、クエリ・キー・バリューを作る
- で、すべての単語の組み合わせの点数の表(3行 × 3列)を作る
- (マスクありのとき)自分より後ろの位置の点数を にする
- 表の 行ごと にソフトマックスをとり、注目の重みにする
- 注目の重み @ で、各単語が取り出す情報を計算する
マスクの表
np.tril は、行列の左下(対角線を含む)だけを残す関数です。3単語なら、次の表の True(○)の位置だけを使います。
| 1語目を見る | 2語目を見る | 3語目を見る | |
|---|---|---|---|
| 1語目 | ○ | × | × |
| 2語目 | ○ | ○ | × |
| 3語目 | ○ | ○ | ○ |
np.where(mask, scores, -np.inf) で、× の位置の点数を にします。 なので、その位置の重みは0になります。
行ごとのソフトマックス
点数の表は3行あり、各行が「その単語が、ほかのどの単語に注目するか」です。そのため、ソフトマックスは 行ごと に計算します。NumPy では、axis=1(横方向)に計算し、keepdims=True で形を保ったまま割ります。
e = np.exp(scores - scores.max(axis=1, keepdims=True)) # 行ごとに最大値を引く
w = e / e.sum(axis=1, keepdims=True) # 行ごとの合計で割る → 各行の合計が1
keepdims=True がないと、行ごとの合計が「3個の数の列」ではなく「3個の数の並び」になり、割り算の向きがずれてしまいます(第11章のミニプロジェクトのブロードキャストと同じ注意です)。
実行の結果(演習の例)
| マスクなしの注目の重み | マスクありの注目の重み | |
|---|---|---|
| 1語目 | 0.401・0.198・0.401 | 1.000・0・0 |
| 2語目 | 0.164・0.164・0.673 | 0.5・0.5・0 |
| 3語目 | 0.178・0.088・0.734 | 0.178・0.088・0.734 |
マスクありでは、1語目は自分だけに、2語目は1語目と自分だけに注目しています。3語目は、もともと全部の単語が「自分より前か自分」なので、マスクの有無で変わりません。どの行も、合計は1です。