本文へスキップ
ひもとくAI

Transformer:Self-Attention・マルチヘッド・位置エンコーディング

レッスン 5/7

Transformer:Attention だけで、全部の位置を一度に見る

RNN の弱点:1つずつ順番にしか計算できない

RNN は、時刻 tt の隠れ状態を計算するのに、時刻 t−1t-1 の結果が必要です。1000語の文なら、1000回の計算を順番に行うしかなく、コンピューターが得意な「たくさんの計算を同時に行う」ことができません。また、離れた位置の情報は、間のすべての時刻を通って伝わるので、薄れやすくなります。

2017年に発表された Transformer は、RNN を使わず、Attention だけで 系列を扱うモデルです。今の大規模言語モデルの多くが、この Transformer を土台にしています。

Self-Attention:文の中の単語どうしで注目し合う

前のレッスンの Attention では、デコーダ(クエリ)がエンコーダ(キー・バリュー)に注目しました。このように、別の系列に注目する Attention を Source-Target Attention(クロスアテンション)と呼びます。Self-Attention(自己注意)では、同じ文の中の単語どうし が注目し合います。各単語について、

  • クエリ qq:「自分は、どんな情報を探しているか」
  • キー kk:「自分は、どんな情報を持っているか」
  • バリュー vv:「注目されたときに渡す情報」

の3つを、単語の埋め込みベクトルに、それぞれ別の重みの行列(WQW_Q、WKW_K、WVW_V)を掛けて作ります。そして、すべての単語の組み合わせで、前のレッスンの3段階(内積 → ソフトマックス → 重み付きの平均)を行います。

例えば「その 犬 は 疲れて いた ので 眠った」の「眠った」は、「犬」に強く注目することで、「誰が眠ったのか」の情報を取り込めます。何語離れていても、1回の計算で直接注目できるのが、RNN との大きな違いです。

まとめて行列で計算する

すべての単語のクエリを並べた表を QQ、キーの表を KK、バリューの表を VV とすると、Self-Attention は次の1行で書けます。

Attention(Q,K,V)=softmax ⁣(QK⊤d)V\text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{Q K^{\top}}{\sqrt{d}}\right) V

QK⊤Q K^{\top} は、「単語の数 × 単語の数」の、すべての組み合わせの内積の表です。ソフトマックスは、表の 行ごと(各単語が、ほかのどの単語に注目するか)に計算します。全部の位置を同時に計算できるので、たくさんの計算を並べて行う GPU で、とても速く動きます。

なぜ d\sqrt{d} で割るのか

dd は、クエリとキーのベクトルの長さ(数の個数)です。数の個数が多いほど、内積は多くの掛け算の合計になるので、値のばらつきが大きくなります。8つのキーとの内積をソフトマックスに通し、一番大きい重みの平均を測りました(クエリとキーは、でたらめな数)。

ベクトルの長さ dd41664256
d\sqrt{d} で割らない0.5490.7490.8580.928
d\sqrt{d} で割る0.3590.3560.3580.368

割らないと、dd が大きいほど、ソフトマックスが1か所にほぼ全部の重みを集めてしまいます(256では平均0.928)。そうなると、ほかの位置への勾配がほとんど0になり、学習が進みにくくなります(第8章の勾配消失と同じ問題)。d\sqrt{d} で割ると、dd によらず同じくらいの広がりに保てます。この式を scaled dot-product attention と呼びます。

Transformer の部品:マルチヘッド・位置・マスク

マルチヘッド Attention:いろいろな見方で、同時に注目する

1つの Self-Attention では、各単語は1通りの注目の仕方しかできません。しかし、文の中の関係には、「誰が(主語)」「何を(目的語)」「どれを指すか(代名詞)」など、いろいろな種類があります。

そこで、クエリ・キー・バリューを作る重みを何組か用意し、それぞれで別々に Attention を計算して、結果を並べてつなげます。これを マルチヘッド Attention(Multi-Head Attention)と呼び、1組1組を ヘッド と呼びます。ヘッドごとに、違う種類の関係に注目するように学習されることが期待されます。

位置エンコーディング:順番の情報を足す

Self-Attention は、すべての組み合わせの内積を計算するだけなので、単語の順番を入れ替えても、各単語の結果は変わりません。これでは「ねこ が いぬ を」と「いぬ が ねこ を」を区別できません。

そこで、各単語の埋め込みベクトルに、「何番目の位置か」を表すベクトルを足します。これを 位置エンコーディング と呼びます。元の Transformer では、位置 pospos ごとに、周期の違う sin と cos の値を並べました。例えば、4個の数のベクトルでは次のようになります。

位置sin⁡(pos)\sin(pos)cos⁡(pos)\cos(pos)sin⁡(pos/100)\sin(pos / 100)cos⁡(pos/100)\cos(pos / 100)
00.0001.0000.0001.000
10.8410.5400.0101.000
20.909−0.4160.0201.000
30.141−0.9900.0301.000

時計の秒針・分針・時針のように、速く回る数と、ゆっくり回る数を組み合わせることで、どの位置も違うベクトルになります。位置のベクトルも学習で決める方法もあります。

マスク:未来の単語を見ないようにする

文章を先頭から1語ずつ生成するモデルでは、学習のときに、ある位置の単語が それより後ろの単語に注目してはいけません(本番では、後ろの単語はまだないため)。そこで、ソフトマックスの前に、後ろの位置の内積を −∞-\infty(マイナス無限大)にします。e−∞=0e^{-\infty} = 0 なので、後ろの位置への注目の重みは0になります。これを マスク(因果マスク)と呼びます。

そのほかの部品

  • フィードフォワード層: Attention の後に、各位置で別々に、第6章の2層のニューラルネットワークを通す
  • 残差接続: 層の入力を、層の出力にそのまま足す。足した道を通って勾配がそのまま届くので、深く重ねても学習しやすくなる(画像のモデルの ResNet で広まった工夫)
  • 層正規化(Layer Normalization): 第8章の Batch Normalization に似ているが、ミニバッチ方向ではなく、1つのデータの中の数の平均と分散でそろえる。系列の長さやミニバッチの大きさに左右されない

Transformer は、「マルチヘッド Attention → フィードフォワード層」(それぞれ残差接続と層正規化付き)のまとまりを、何段も重ねたモデルです。

BERT と GPT、PyTorch、振り返り

Transformer から生まれたモデル

元の Transformer は、翻訳のためのエンコーダ・デコーダの形でした。その後、どちらか一方だけを使うモデルが広まりました。

形代表例学習の仕方得意なこと
エンコーダだけBERT文の一部の単語を隠し、前後の両方から当てる(マスク付き言語モデル)文の分類、質問に対する答えの場所の抽出など、文を「読む」こと
デコーダだけGPT前の単語から、次の単語を当てる(因果マスクを使う)文章の生成、会話
エンコーダ・デコーダ元の Transformer入力の文から出力の文を生成する翻訳、要約

どれも、大量の文章で、まず「単語を当てる」練習(事前学習)をしてから、目的の課題に合わせて少しだけ学習し直す(ファインチューニング)使い方が基本です。正解付きのデータが少なくても、事前学習で身に付けた言葉の知識を生かせます。

画像も、小さな区画に分けて「単語」のように並べれば、Transformer で扱えます(Vision Transformer)。

PyTorch で Attention を使う

import torch
import torch.nn.functional as F
from torch import nn

Q = torch.randn(1, 5, 8)    # (データの数, 単語の数, ベクトルの長さ)
out = F.scaled_dot_product_attention(Q, Q, Q, is_causal=True)   # 因果マスク付きの Self-Attention

mha = nn.MultiheadAttention(embed_dim=8, num_heads=2, batch_first=True)   # 2つのヘッド
out, weights = mha(Q, Q, Q)
print(weights.shape)        # → torch.Size([1, 5, 5])(各単語が、どの単語に注目したか)

F.scaled_dot_product_attention は、スライドの式 softmax(QK⊤/d)V\text{softmax}(QK^{\top} / \sqrt{d})V そのものです。ミニプロジェクトで NumPy で書くものと同じ値になることを、運営側で確かめています。

よくある誤解

  • 「Transformer は、単語の順番を自然に理解している」: Self-Attention だけでは順番が分からないので、位置エンコーディングで順番の情報を足しています
  • 「d\sqrt{d} で割るのは、計算を速くするため」: 内積のばらつきを抑えて、ソフトマックスが1か所に偏り、勾配が消えるのを防ぐためです
  • 「Transformer は、どんなに長い文でも同じ速さで処理できる」: すべての単語の組み合わせの内積を計算するので、単語の数が2倍になると、計算の量は約4倍になります。長い文を扱うための工夫が、今も研究されています

振り返り

  • Transformer は、RNN を使わず、Self-Attention で全部の位置を同時に見る。softmax(QK⊤/d)V\text{softmax}(QK^{\top} / \sqrt{d})V
  • d\sqrt{d} で割るのは、ソフトマックスが偏って勾配が消えるのを防ぐため
  • マルチヘッドでいろいろな関係に注目し、位置エンコーディングで順番を足し、マスクで未来を隠す。残差接続と層正規化で深く重ねる
  • エンコーダだけの BERT、デコーダだけの GPT。事前学習とファインチューニング

演習

演習を読み込んでいます…