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

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

トークンをもう一つ作るとき何を計算し直すのか

TT Labで続きを見る

一言でいうと

自己回帰生成は、トークンを一度に1つずつ作ります。キャッシュがないと、トークンを1つ作り足すたびに、前のコンテキスト全体のk・vを最初から作り直します。KVキャッシュは、その作り直しをなくします。式はそのままで、計算の回数だけを変えます。

なぜ必要なのか

アテンション自体は、すでに手で計算してみたことでしょう。スコアを出し、ソフトマックスをかけ、値の加重平均をとります。3行です。

問題は、その3行を生成のループの中で呼ぶときに生じます。長さ100の答えを作るには、モデルを100回呼びます。1回呼ぶたびに、モデルは「ここまでのコンテキスト全体」を入力として受け取ります。そして、そのコンテキストのすべての位置について、q・k・vを作ります。

ところが、2回目の呼び出しで作る最初のトークンのk・vは、1回目の呼び出しで作ったものとまったく同じ値です。3回目の呼び出しでも同じです。100回目の呼び出しでも同じです。同じ値を100回作って、99回捨てているのです。

この無駄は、コードを読んでもなかなか見えません。アテンションの関数には、間違っているところが1つもないからです。見えるようにするには、回数を数える必要があります。

なぜ時間ではなく乗算を数えるのか

「キャッシュをオンにしたら速くなった」という測定からは、何も学べません。速くなった度合いは、マシン、負荷、バッチサイズ、メモリの帯域幅によって変わり、同じマシンでも、測り直すと違う数が出ます。

乗算の回数は違います。同じコードを同じ入力で動かせば、いつも同じ数です。そして、その数は、長さに対してどう増えていくかをそのまま示します。そのため、このラボは時間を測りません。その代わりに、乗算を行う関数を1つ作って、その中で数えます。

def mul(a, b):
    global _MULS
    _MULS += 1
    return a * b

正直な方法です。乗算が起きる箇所をすべてこの関数に通せば、誰がどこで何回掛けたかを数えられます。計算量を推測せずに、数えるのです。

数えてみると出てくる形

1つの層、1つのヘッドだけを置いて、モデルの次元をd、いまのコンテキスト長をnとします。q・k・vを作る重みは、d×dの行列3つです。

キャッシュなしで、トークン1つ:

キャッシュありで、トークン1つ:

違いは、最初の行1つだけです。キャッシュなしでは、射影のコストがコンテキスト長に比例して増え、キャッシュありでは、そのコストが長さと無関係に固定されます。アテンション自体(後ろの2行)は、どちらも長さに比例して増えます。キャッシュは、それをなくしてはくれません。キャッシュがなくすのは、再計算であって、アテンションではありません。

トークンを5回作る間に、各位置のキーと値を何回作るかを数えた図。キャッシュがないと、毎回前の位置をすべて作り直すので、1、2、3、4、5回となり、合計15回で、キャッシュを使うと、新しく付いた位置1つだけを作るので、毎回1回、合計5回です

このラボでは、その数を手で数えて、長さごとに表を作ります。他の人が書いておいた倍率を書き写すのではなく、自分で数えた数を書きます。

キャッシュが成り立つ理由は、因果マスク

なぜ、前の位置のk・vをそのまま再利用してよいのでしょうか。新しいトークンが付いたのに、前の位置の値が変わらないと、どうして確信できるのでしょうか。

因果マスクがあるからです。Attention Is All You Needのデコーダーは、各位置が自分より前だけを見るようになっています。位置3のkとvは、位置3の入力だけで作られ、4番目のトークンが後ろに付いても、位置3の値が変わる理由がありません。

もしマスクがなくて、すべての位置が互いを見るなら、キャッシュは成り立ちません。後ろに何かが付くたびに、前の表現が変わるからです。そのため、KVキャッシュはデコーダー専用の構造の性質であって、どこにでも使える最適化ではありません。

ここで、重要な結論が1つ出てきます。キャッシュは近似ではありません。答えは変わりません。同じ値を再び作らないだけです。ですから、キャッシュをオンにしたときとオフにしたときで出力が変わるなら、それはキャッシュの性質ではなく、実装の欠陥です。

キャッシュが食うメモリ

節約するものがあれば、支払うものもあります。キャッシュは、メモリを食います。要素の数は、積で書かれます。

원소 수 = 2 (K 와 V) × 층 수 × 헤드 수 × 문맥 길이 × 헤드 차원
바이트  = 원소 수 × 자료형 한 원소의 바이트 수

このコードブロックの韓国語の式は、要素数が2(KとV)×層の数×ヘッドの数×コンテキスト長×ヘッド次元で、バイト数は要素数×データ型1要素のバイト数だ、という意味です。

ここで注目したいのは、コンテキスト長が掛けられるという点です。長さが2倍なら、キャッシュも2倍です。これが長いコンテキストが高くつく理由の1つで、1台のマシンで同時に処理できるリクエスト数を決めるのも、たいていはこの表です。

このラボでは、層・ヘッド・ヘッド次元を自分が決めた小さな値にして、その値だけで計算します。実際のモデルのGBの数字は使いません。ここで測っていないからです。測っていない数字を書き写した瞬間に、その文章は根拠を失います。データ型を変えると最後の項だけが変わること、長さが掛けられること、この2つは、自分で作った表の中で直接見えます。

現場での姿

第1に、長いプロンプトを毎回送り直して、キャッシュを使えません。会話を続けるときに、前の内容をまるごと送り直すと、サーバー側でキャッシュを引き継いで使う根拠がありません。Hugging FaceのCache strategiesのドキュメントが、キャッシュのオブジェクトを呼び出しの間に自分で持ち回る方法を別に説明している理由が、これです。

第2に、同時リクエスト数の上限が、計算ではなくメモリから来ます。キャッシュは、リクエストごとに別々に確保され、長さに比例して増えます。そのため、長い会話の数件が、短いリクエストの数十件より多くの場所を占めます。

第3に、キャッシュを低い精度で保持すると、出力がわずかに変わります。キャッシュ自体は近似ではありませんが、キャッシュにより不正確に書き込むことは近似です。このラボでは、キャッシュに入れる前に丸めてみて、その差が許容誤差を超えるかどうかを、自分で確認します。

第4に、最初のトークンと、その次のトークンでは、性格が違います。プロンプトをまるごと読む最初の計算は、キャッシュが空なので、節約できるものがありません。キャッシュが得をするのは、2番目のトークンからです。Hugging FaceのText generationのドキュメントも、生成をこの2つのフェーズに分けて説明しています。

第5に、キャッシュを使いながらバッチを混ぜると、位置がずれます。キャッシュの行の順序が、そのままトークンの順序です。リクエストをまとめたり分けたりしながら、行を間違ってつなげると、エラーは出ず、言葉がおかしくなるだけです。

実務で本当に大切なこと

次のラボですること

/root/work/tf-kv/kv.pyを、1ステップずつ育てていきます。標準ライブラリだけを使います。このPodのシステムのPythonにはnumpyがなく(/opt/onnx-lab/bin/pythonの中にしかありません)、torchもtransformersもありません。そのため、ここに出てくる数字は、すべて自分が作ったコードが数えたものです。

乗算を数える乗算関数と内積から始めて、トークン1つを射影する関数、クエリ1つが積み重なったK・Vを見るアテンション、キャッシュなしで動く版、キャッシュで動く版を、順に作ります。そのあと、2つの版を同じ入力で動かして、出力が同じ値であるかを、許容誤差で確認し、長さごとに乗算の回数を数えて、表にします。

最後の2つのステップが、このラボの要点です。キャッシュに丸めて入れてみると、2つの方式の出力が、もはや許容誤差の範囲に収まらなくなります。キャッシュは近似ではありませんが、キャッシュにより不正確に書き込むことは近似だという事実が、数字で出ます。そして、自分が決めた層・ヘッド・ヘッド次元・データ型で、キャッシュが食う要素の数とバイト数を計算して、乗算の回数の表といっしょに、記録に残します。採点ツールは、自分で作ったモジュールを実際に呼び出し、毎回異なるサイズの入力で関数を直接叩いて、値と乗算の回数を別に計算して照合します。