TT Lab
はじめる
学ぶ 学習パス コース

Transformer — アテンションを手で計算する

番号がベクトルになり、ベクトルが点数になる

TT Labで続きを見る

一言でいうと

埋め込みは行を取り出す表にすぎず、出力側のロジットはその表をもう一度使う乗算です。その間に入っている数の個数が、語彙サイズ掛ける幅で、語彙を大きくしたときの値は、そこで支払います。

なぜ必要なのか

トークン化まで終えると、文章は整数のリストになります。ところが、整数はそのままでは使えません。トークン3とトークン4が隣同士という意味ではないのに、数字のままにしておくと、すべての演算がそう読んでしまいます。そこで、番号ごとにベクトルを1つ付けます。そのベクトルを1枚に積み重ねたものが、埋め込み表です。

ここで最初に引っかかるのが、なぜルックアップなのかです。教科書は「ワンホットベクトルに行列を掛ける」と書くのに、コードはtable[token_id]の1行です。2つが別の話に見えます。違いはありません。ワンホットは、1か所だけが1で、残りが0なので、掛けて足すと、その1行だけが生き残ります。値はまったく同じで、乗算の回数だけが違います。ルックアップは、別の演算ではなく、同じ演算の近道です。

2番目に引っかかるのが、出ていく道です。アテンションとブロックをすべて通ると、ベクトルが1つ残ります。それを再び単語に戻す必要があり、語彙全体にスコアを付ける作業になります。幅dのベクトルを、語彙サイズVのスコアに広げるには、V×dの行列が必要です。ところが、その形の行列は、すでにあります。入力側の埋め込み表が、まさにその形です。

ルックアップは乗算の近道にすぎない

表の形は、(語彙数V、幅d)です。持っている数の個数は、V掛けるdです。幅をそのままにして、語彙だけを2倍にすると、その個数も2倍になります。トークンの数を減らすために語彙を大きくしましたが、その値は、この表から出ていきます。

ids = [7, 7, 41]
rows = [table[i] for i in ids]      # 조회
# 같은 값을 원-핫으로 계산하면
one = [0.0] * V; one[7] = 1.0
row = [sum(one[r] * table[r][c] for r in range(V)) for c in range(d)]
# rows[0] 과 row 는 같은 값이다. 곱셈만 V 곱하기 d 번 더 했다.

ここで、もう1つのことが現れます。idsの前の2つの位置は、同じ番号なので、まったく同じベクトルが出ます。前に何があっても、後ろに何が来ても、同じです。埋め込みには、コンテキストがありません。同じ単語が位置によって違って読まれるのは、アテンションがあとで行うことで、表はただの表です。torch.nn.MultiheadAttentionの最初の引数がembed_dimなのも、そのためです。埋め込みの幅がそのままモデルの幅なので、表で決めたdを、その後ろのすべての層がそのまま受け取ります。

行列の積を実際に使う場面では、NumPyのmatmulのようなものが代わりに計算してくれます。ただし、このラボのPodのシステムのPythonにはnumpyがなく、/opt/onnx-lab/bin/pythonの中にしかないので、ここでは標準ライブラリで直接2つの方法を計算して、値が同じかを見ます。

出口: 同じ表をもう一度使う

Attention Is All You Needは、埋め込みを扱う短い節で、2つのことを書いています。1つは、2つの埋め込み層と、ソフトマックスの前の線形変換が、同じ重み行列を共有して使うということで、もう1つは、埋め込み層で、その重みに√dを掛けるということです。前者が、重み共有(weight tying)です。

共有すると、2つのことが生じます。第1に、表が1式になるので、数の個数が半分になります。別々に置くと、V・dが2式なので、2・V・dです。第2に、スコアを付ける方式が内積になります。隠れ状態のベクトルhに対して、トークンtのスコアは、表のt番目の行とhの内積です。そのため、hがあるトークンの埋め込みと同じになると、そのトークンのスコアが最も大きくなります。自分自身との内積は、長さの2乗なので、他のどんな内積よりも大きくなりやすいからです。

この性質は便利ですが、落とし穴も同じ場所から出てきます。共有すると、1つの表が、入ってくる意味と出ていくスコアという2つの仕事を同時に担います。片方によい配置が、もう片方にもよいという保証はありません。共有するかどうかは、そのため、無料の選択ではなく、取引です。パラメーターを半分に減らす代わりに、表に2つの仕事をさせます。

近いトークンは内積で探さない

「このトークンに近いトークン」を探すとき、内積をそのまま使ってはいけません。内積には、相手の長さが掛けられて入っているからです。方向が少し合っていなくても、長ければ、前に出てきます。

コサインは、その長さを割って消します。方向だけを残すのです。違いを確認する最も確実な方法は、表の1行だけを何倍かに引き伸ばしてみることです。その行のコサインはまったく変わらず(方向がそのままです)、内積はすべて、その倍数の分だけ大きくなります。近傍のリストをコサインで取り出すと順序がそのままなのに、内積で取り出すと、引き伸ばした行が先頭に飛び出します。

そのため、ロジットの順序と「意味が近い順」は、同じものではありません。ロジットは内積の順序で、そこには長さが混ざっています。

論文が√dを掛ける箇所

同じ節のもう1つの文が、埋め込みに√dを掛けるというものです。掛けると、何が変わるのでしょうか。方向はまったく変わりません。すべてのマスに同じ数を掛けたので、コサインはそのままです。変わるのは大きさだけで、大きさは、ちょうど√d倍になります。

大きさがなぜ重要なのでしょうか。埋め込みに位置情報を足す場面で、2つの信号の大きさがあまりに違うと、片方が埋もれます。大きさをそろえておく作業が、そのために必要です。このラボでは、「なぜよりによって√dなのか」を証明する代わりに、掛けると大きさが√d倍になり、方向はそのままであることを、数字で確認します。測って言えるのは、そこまでです。

現場での姿

第1に、語彙を大きくしようという提案が、メモリの会議で終わります。トークンが減ってよさそうだと思って、語彙を2倍にしようと言うと、表が2倍になるという答えが返ってきます。幅には手を付けていないのに、そうなります。どちらが得かは、2つの値を並べて数えてみる必要があります。

第2に、「共有したか、しなかったか」で、パラメーターの数がずれます。同じ構成を書いておいたのに、計算した値が、表1つ分の差になります。表が1つまるごとあるかないかなので、丸め誤差のようなものではありません。

第3に、類似トークンのリストが変です。内積で取り出しておいて「意味が近い」と呼ぶと、長さの大きい行が、どのクエリにも割り込んできます。コサインに変えると、その行が消えます。

第4に、埋め込みだけを取り出して、コンテキストを期待します。同じ単語は、表の中ではいつも同じ行です。文を表すベクトルが必要なら、モデルを通す必要があり、表をルックアップした値は、コンテキストのない値です。

第5に、幅を変えると、その後ろがすべて連動して動きます。埋め込みのdは、その後ろのすべての層が受け取る幅なので、1か所だけを直すことはできません。

実務で本当に大切なこと

次のラボですること

/root/work/tf-embed/embed.pyを、1ステップずつ育てていきます。標準ライブラリだけを使います。このPodのシステムのPythonにはnumpy・torch・transformersがなく、numpyは/opt/onnx-lab/bin/pythonの中にしかありません。実際のモデルを呼び出さないので、実際のモデルの語彙サイズやパラメーターの数のような数字は使いません。ここに出てくる値は、すべて自分が作った表から測ったものです。

決定的な埋め込み表を作ることから始めて、形とパラメーターの数を数え、番号で行を取り出し、同じ値をワンホットの積で計算し直して、2つの値が同じかを見ます。そのあと、同じ表でロジットを出し、共有したときと別々に置いたときの数の個数を数え、語彙を2倍にしてみます。

後ろの2つのステップが要点です。表の1行だけを何倍かに引き伸ばして、近傍のリストを、コサインと内積でそれぞれ取り出します。コサインのリストはそのままなのに、内積のリストでは、引き伸ばした行が先頭に出ます。最後に、埋め込みに√dを掛けて、大きさがちょうど√d倍になることと、方向がそのままであることを、並べて測ります。採点ツールは、自分で作ったモジュールを実際に呼び出し、毎回異なる表と異なる番号で関数を直接叩いて、採点ツールが別に計算した値と照合します。