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

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

文脈を倍にすると何が四倍になるのか

TT Labで続きを見る

一言でいうと

Transformerの1層の計算は、2つの取り分に分かれます。アテンションのスコアと混合は、コンテキスト長の2乗で増え、残りの線形変換は、長さに比例して増えます。短いコンテキストでは後者が支配し、ある長さを超えると前者が支配します。その位置を知ることが、この文章のすべてです。

なぜ必要なのか

コンテキストの上限を延ばしてほしいという要望は、どのチームにも来ます。4千から8千へ、8千から3万2千へ。延ばすとどれだけ高くなるのかという問いに、「2倍くらい」と答えて、請求書を見て驚くことが繰り返されます。

逆の場合もあります。コンテキストを2倍に延ばしたのに、コストが2倍しか増えず、「2乗だと聞いたのに、違うじゃないか」と見過ごしてしまうケースです。どちらも、同じ誤解から生じます。どちらか片方しか見ていないのです。

1つの層の中には、性質の異なる2つの計算が混ざっています。アテンションは、位置ごとにコンテキスト全体を走査するので、位置の数が増えると、走査する作業まで一緒に増えて、2乗になります。一方、クエリ・キー・値を作る射影、出力の射影、フィードフォワードは、位置1つ1つに同じように1回ずつ行う作業なので、位置の数だけに比例します。

そのため、「コンテキスト長に対して2乗」という言葉は、半分しか合っていません。正確には、2乗で増える取り分と、比例して増える取り分の和であり、どちらが支配するかは、長さとモデルの幅が一緒に決めます。

何を数えればよいのか

ここでは、時間を測ってはいけません。同じコードでも、マシンと負荷によって違い、コンテナの中ではさらに揺れます。その代わりに、乗算の回数を数えます。回数は、どこで動かしても同じ整数です。

数え方は単純です。長さdの2つのベクトルの内積は、乗算がちょうどd回です。あとは、その内積が何回入るかを数えるだけです。

def dot(a, b, ctr):
    total = 0.0
    for x, y in zip(a, b):
        total += x * y
    ctr.add(len(a))      # 곱셈 d 번을 장부에 적는다
    return total

アテンションの定義を、この物差しで測ってみると、次のようになります。

そのため、2乗で増える取り分は2 * n * n * dです。反対側は、位置1つが支払う分を1回だけ数えて、位置の数を掛ければ済みます。

合わせて、位置ごとに4 * d * d + 2 * d * d_ffで、位置がn個あります。

交差点という考え方

2つの式を並べると、交差点が手で解けます。2 * n * n * dがn * (4 * d * d + 2 * d * d_ff)以上になる、最も小さいnを探せばよいのです。両辺を2 * n * dで割ると、条件はn >= 2 * d + d_ffに絞られます。

1つの層の乗算の回数を、2つの取り分に分けて描いた図。アテンションのスコアと混合は、コンテキスト長の2乗で増える曲線で、残りの線形変換は、長さに比例する直線です。2本の線は、nが2d足すd_ffの位置で交わり、その手前では線形の取り分が、後ろでは2乗の取り分が支配します

数字が与える感覚が重要です。幅の広いモデルほど、交差点が後ろへずれます。位置1つが支払う線形のコストは、幅の2乗で大きくなるのに、アテンション側は、幅に比例してしか大きくならないからです。大きなモデルで「2乗なのに、なぜ感じないのか」という体感が生まれるのは、ここからです。まだ交差点の手前にいるのです。

そして、交差点を過ぎると、話が変わります。長さを2倍にすると、合計が4倍に近づき始めます。

因果マスクが捨てる半分

デコーダーのアテンションには、因果マスクがあります。i番目の位置は、自分自身までしか見ないので、行の長さが1、2、3と増えて、nで終わります。実際に使われるスコアは、n * (n + 1) / 2個です。

行列は、依然としてnの2乗のマスです。つまり、半分近くは計算しておいて捨てているのです。正確には、(n - 1) / (2 * n)が捨てられ、長さが長くなるほど、半分に近づきます。

素朴に全体を計算してマスクをかぶせる実装が、まさにそうしています。計算を減らすには、マスクをあとからかぶせるのではなく、最初から計算しない必要があり、そのために、ブロック単位で三角形だけを回すカーネルが登場しました。PyTorchのscaled_dot_product_attentionがis_causalを別に受け取るのも、同じ理由です。マスクをテンソルで受け取って掛けることと、「因果です」と伝えることは、内部で動く処理が異なります。

メモリも同じ形

スコア行列をまるごと保持すると、値の個数はヘッドごとにnの2乗です。合計でn * n * h個で、これにデータ型1つ分のバイト数を掛けると、バイト数が出ます。

これがなぜ痛いのかは、逆から見るとわかります。バジェットを決めておいて、入る最大の長さを求めると、バジェットを4倍に増やして、ようやく長さが2倍になります。GPUを2倍に増やしても、コンテキストは1.41倍しか延ばせないということです。

そのため、スコア行列をまるごと作らない実装が重要になりました。ブロックに切って、1片ずつ処理して捨てれば、同じ計算を行いながらも、保持する値の数が減ります。計算量はそのままで、メモリだけが減るのです。この2つが別々に動くという感覚がないと、「なぜ計算は同じなのに、より長いコンテキストになるのか」が理解できません。

入れる長さと作る長さは、増え方が違う

プロンプトを一度に入れる作業(プリフィル)と、トークンを1つずつ作る作業(デコード)は、増え方の形が違います。

プリフィルは、n個の位置を一度に処理するので、スコアがnの2乗の規模で生じます。一方、すでにn個が積み重なったあとに、もう1個を足すときは、新しいクエリ1つがn + 1個のキーを見るだけです。行が1つです。そのため、g個を作る間に生じるスコアの総和は、g * n + g * (g + 1) / 2になります。

ここで、おもしろい恒等式が1つ出てきます。nを0にすると、この値はcausal_pairs(g)とちょうど同じになります。一度に入れてマスクで消しても、1つずつ足しても、実際に必要なスコアの数は同じです。分かれるのは、それを一度に行うか、分けて行うかだけです。その差を何で埋めるのかは、KVキャッシュのドキュメントが扱うテーマで、次のモジュールの担当です。

現場での姿

第1に、コンテキストの上限を2倍に延ばしたのに、遅延が2倍に増えず、安心します。まだ交差点の手前なので、線形の取り分が支配しているだけです。長さをさらに延ばすと、急に傾きが変わります。

第2に、同じ長さなのに、モデルを変えたらコスト曲線の形が変わります。幅が違えば、交差点は別の場所にあります。あるモデルで測って得た倍率を、別のモデルにそのまま使うと、ずれます。

第3に、長いコンテキストで、メモリが先にあふれます。計算は持ちこたえるのに、スコア行列を置く場所がありません。計算量とメモリが、同じnの2乗であっても、先にぶつかる壁は、たいていメモリのほうです。

第4に、マスクをテンソルにして掛ける実装が、2倍遅くなります。全体を計算して、半分を捨てるからです。同じ式でも、いつマスクを使うかによって、実際の仕事量が分かれます。

第5に、プリフィルは遅いのに、トークン生成は速い、あるいはその逆になります。この2つは別の形で増えるので、1つの倍率でまとめて見積もると、必ずどちらかが外れます。

実務で本当に大切なこと

次のラボですること

/root/work/tf-cost/cost.pyを、1ステップずつ育てていきます。標準ライブラリだけを使います。このPodのシステムのPythonにはnumpyがなく(/opt/onnx-lab/bin/pythonの中にしかありません)、torchもtransformersもありません。その代わりに、乗算を数えるカウンターを手で作り、そこから出た整数だけを使います。

カウンターと内積から始めて、スコア行列と混合を実際に動かして、nの2乗が出ることを、数えて確認します。そのあと、位置1つが支払う線形変換6つを実際に動かして、そちらが長さと無関係であることを確認し、2つの取り分を並べた表を作ります。

その表から、交差点を探します。モデルの幅を変えながら、交差点がどこへ動くかも、自分の数字で見ます。続いて、因果マスクが残すスコアを数えて、半分が捨てられていることを確認し、スコア行列をまるごと保持するときのバイト数と、バジェットに収まる最大の長さを求めます。

最後のステップが、このラボの要点です。1トークンを足すときに新しく生じるスコアは、行1つだけであることを数え、長さ0からg個を作るときの総和が、因果マスクが残した数とちょうど同じになるという恒等式を、整数で確認します。採点ツールは、自分で作ったモジュールを実際に呼び出し、毎回異なるサイズで関数を直接叩いて、カウンターが実際にどれだけ上がったかまで照合します。値を暗記して入れることはできません。