文脈を倍にすると何が四倍になるのか
一言でいうと
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
アテンションの定義を、この物差しで測ってみると、次のようになります。
- スコア行列QK: 行もn個、列もn個で、列ごとに長さdの内積 →
n * n * d - 混合AV: 出てくるのはn×dと小さいですが、位置ごとにコンテキスト全体のn個を走査します → またも
n * n * d - ヘッドをh個に分けても、合計は同じです。ヘッドごとに幅が
d / hに減り、そのようなヘッドがh個あるからです。
そのため、2乗で増える取り分は2 * n * n * dです。反対側は、位置1つが支払う分を1回だけ数えて、位置の数を掛ければ済みます。
- クエリ・キー・値の射影3つと、出力の射影1つ → 位置ごとに
4 * d * d - フィードフォワードの2層 → 位置ごとに
2 * d * d_ff
合わせて、位置ごとに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倍にすると、合計が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つの倍率でまとめて見積もると、必ずどちらかが外れます。
実務で本当に大切なこと
- 「2乗」という言葉だけを覚えず、2つの取り分に分けて見てください。どちらが支配するかは、長さと幅が一緒に決めます。
- 時間の代わりに、回数を数えてください。回数は、マシンが変わっても同じで、見積もりを他の人に説明するときの根拠になります。
- 交差点を、自分のモデルの数字で求めておいてください。その位置の前後で、容量計画が変わります。
- 計算量とメモリを、別々に見てください。同じnの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個を作るときの総和が、因果マスクが残した数とちょうど同じになるという恒等式を、整数で確認します。採点ツールは、自分で作ったモジュールを実際に呼び出し、毎回異なるサイズで関数を直接叩いて、カウンターが実際にどれだけ上がったかまで照合します。値を暗記して入れることはできません。