文脈長のコストを自分で数える
目標
乗算を数えるカウンターを作り、コンテキスト長nが増えるとき、何がnの2乗で増え、何がnに比例して増えるのかを、自分で数えて確認します。2つの取り分を並べた表を作り、nの2乗の取り分が線形の取り分に追いつく交差点を、自分で数えた数字から探します。因果マスクが残すスコアがn(n+1)/2個だけであること、スコア行列をまるごと保持すると要素がnの2乗個になること、1トークンを足すときに新しく生じるスコアは行1つだけであることまで、整数で数えます。
なぜ重要なのか
「コンテキスト長に対して2乗」という言葉は、半分しか合っていません。1つの層の中には、性質の異なる2つの計算が混ざっています。アテンションのスコアと混合は、位置ごとにコンテキスト全体を走査するので、2乗で増えますが、クエリ・キー・値の射影と、出力の射影と、フィードフォワードは、位置1つ1つに同じように1回ずつ行う作業なので、位置の数だけに比例します。どちらが支配するかは、長さとモデルの幅が一緒に決めます。そのため、コンテキストの上限を延ばしたときのコストを1つの倍率で見積もると、必ず外れます。
このラボは、時間を測りません。同じコードでも、マシンと負荷によって違い、コンテナの中ではさらに揺れるので、測っても比べられません。その代わりに、乗算の回数を数えます。回数は、どこで動かしても同じ整数なので、他の人に根拠として見せられます。ここに出てくる数字はすべて、自分のカウンターが数えた値で、実際のモデルの秒数やGBは測っていないので、使いません。
モデルの形は、Attention Is All You Needのbase設定を仮定として使います。D_MODEL = 512、D_FF = 2048、N_HEADS = 8です。これはこのラボが決めた仮定であって、自分が使うモデルの値ではありません。
採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なるサイズで関数を直接叩いて、カウンターが実際にどれだけ上がったかまで照合します。サイズは実行のたびに変わるので、値を暗記して入れることはできません。
ステップ
- /root/work/tf-cost/cost.pyに、定数
D_MODEL = 512・D_FF = 2048・N_HEADS = 8と、カウンターMulCount、そしてそれを使うdot(a, b, ctr)を作成してください。 attn_scores(Q, K, ctr)・attn_mix(A, V, ctr)・quad_mults(n, d_model)を追加して、nの2乗で増える取り分を数えるようにしてください。matvec(M, x, ctr)・per_position_mults(d_model, d_ff, ctr)・linear_mults(n, d_model, d_ff)を追加して、位置の数に比例する取り分を数えるようにしてください。cost_table(ns, d_model, d_ff)を作成して、長さごとに(n, n제곱 몫, 선형 몫, 합)(プレースホルダーは長さ、2乗の取り分、線形の取り分、合計です)を返すようにしてください。crossover_n(d_model, d_ff)を作成して、nの2乗の取り分が線形の取り分に初めて追いつく長さを探すようにしてください。causal_pairs(n)とwasted_pairs(n)を作成して、因果マスクが残すスコアと、捨てるマスを数えるようにしてください。score_bytes(n, n_heads, itemsize)・max_context_for_bytes(budget_bytes, n_heads, itemsize)・append_pairs(n)・generate_pairs(n, g)を作成してください。- 上の関数を実際に動かした結果を、/root/work/tf-cost/cost_report.jsonと、/root/work/tf-cost/cost_report.mdに記録してください。
参考
- 実行の契約: 採点ツールは、
/root/work/tf-cost/cost.pyをPythonモジュールとして読み込み、D_MODEL・D_FF・N_HEADS・MulCount・dot・attn_scores・attn_mix・quad_mults・matvec・per_position_mults・linear_mults・cost_table・crossover_n・causal_pairs・wasted_pairs・score_bytes・max_context_for_bytes・append_pairs・generate_pairsを直接使います。スクリプトとして実行しないので、if __name__ == "__main__"はなくてかまいません。 MulCountは、mults(これまでに数えた乗算の回数)とcalls(帳簿に書き込んだ回数)の2つの値を持ち、add(k)でkを足します。最初は、どちらも0です。dot(a, b, ctr)は、内積の値を返し、カウンターを正確に長さの分だけ上げます。乗算1回を1と数えるのであって、呼び出し1回を1と数えるのではありません。attn_scores(Q, K, ctr)は、Qの行数×Kの行数のサイズの行列を返します。QとKがそれぞれn×dなら、カウンターはn * n * dだけ上がります。attn_mix(A, V, ctr)は、Aがn×n、Vがn×dのとき、n×dを返し、カウンターはまたもn * n * dだけ上がります。出力が小さいからといって、計算が小さいわけではありません。quad_mults(n, d_model)は、2つの取り分を合わせた2 * n * n * d_modelです。ヘッドの数は入りません。ヘッドごとに幅がd_model / hに減り、そのようなヘッドがh個あるので、合計が同じだからです。per_position_mults(d_model, d_ff, ctr)は、線形写像6つをそれぞれ実際に動かして、乗算を数えます。クエリ・キー・値の射影3つ(d_model x d_model)、出力の射影1つ(d_model x d_model)、フィードフォワードの1層目(d_ff x d_model)と2層目(d_model x d_ff)です。返す値は、この関数が上げた分で、採点ツールは、カウンターが少なくとも6回以上呼ばれたかも見ます。式1つを一度に足して終わりにすると、不合格になります。linear_mults(n, d_model, d_ff)は、位置1つの値に、位置の数を掛けた整数です。カウンターは受け取りません。cost_table(ns, d_model, d_ff)が返す行は、(n, quad, linear, quad + linear)の4つの欄で、順序はnsと同じです。crossover_n(d_model, d_ff)は、quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff)が初めて真になるnです。等号を含みます。1から上げながら探してもよく、式を解いてもかまいません。causal_pairs(n)は、i番目の行がi + 1個を使うという事実から出てきます。causal_pairs(0)は0、causal_pairs(1)は1です。wasted_pairs(n)は、n * nから、実際に使う取り分を引いた値です。score_bytes(n, n_heads, itemsize)は、n * n * n_heads * itemsizeです。ヘッドの数を抜かさないでください。max_context_for_bytes(budget_bytes, n_heads, itemsize)は、score_bytes(n, ...) <= budget_bytesを満たす最大のnです。実数の平方根を四捨五入すると、1つ超えてしまうことがあるので、math.isqrtを使ってください。append_pairs(n)はn + 1、generate_pairs(n, g)は、行の長さn+1からn+gまでの合計です。generate_pairs(0, N)がcausal_pairs(N)と同じである必要があります。- ステップ8のレポートは、
D_MODEL・D_FF・N_HEADSと、ITEMSIZE = 2(2バイトのデータ型を仮定)、バジェット1073741824(1 GiB)、表の長さ[128, 256, 512, 1024, 2048, 4096]、因果とメモリの計算の基準の長さ2048、生成するトークンの数256を使います。 quad_ratio・linear_ratioは、表の最後の2行(2048と4096)の間の倍率です。割り算なので実数で、採点ツールはabs(a - b) <= atol + rtol * abs(b)で比べます。- このPodにはインターネットがありません。
pip installはできず、システムのPythonにはnumpy・torch・transformersがありません。numpyは/opt/onnx-lab/bin/pythonの中にしかありません。標準ライブラリだけで十分です。 - 公式ドキュメント: Attention Is All You Need・PyTorch — scaled_dot_product_attention・Hugging Face — Cache strategies
- よくある間違い: カウンターを呼び出し回数で数える、
attn_mixが出力のサイズの分だけを数えると考える、線形の取り分で出力の射影を抜かす、交差点を等号なしで探して1つずれる、causal_pairsで対角線を除く、score_bytesにヘッドの数を掛けない。
何を測らないのか
時間は測りません。実際のモデルのGBや秒も使いません。測っていない数字を記録に書くと、その記録には根拠がありません。
乗算を数える物差しを作る
/root/work/tf-cost/cost.pyに、定数D_MODEL = 512・D_FF = 2048・N_HEADS = 8と、カウンターMulCount(mults・callsを持ち、add(k)で上げます)、そしてdot(a, b, ctr)を作成してください。dotは、内積の値を返し、カウンターを正確にベクトルの長さの分だけ上げます。
時間を測ろうとしないでください。同じコードでも、マシンと負荷によって変わり、比べられません。長さdの内積は、乗算がちょうどd回なので、ctr.add(len(a))の1行で済みます。呼び出し1回を1と数えると、後ろのすべての数字が崩れます。callsは、帳簿に何回に分けて書き込んだかを数える値で、ステップ3で使われます。
nの2乗で増える取り分を数える
attn_scores(Q, K, ctr)・attn_mix(A, V, ctr)・quad_mults(n, d_model)を追加してください。attn_scoresはn×nのスコア行列を、attn_mixはAでVを混ぜたn×dを返し、どちらもカウンターをn * n * dだけ上げます。quad_multsは、2つを合わせた2 * n * n * d_modelです。
attn_scoresは、クエリごとにキーのすべてとdotすれば終わりです。attn_mixが混乱しやすい箇所です。出てくるのはn×dと小さいですが、出てくる位置ごとにコンテキスト全体のn個を走査する必要があるので、乗算はスコア計算とまったく同じn * n * d回です。Vの縦の列を取り出してdotに渡せば、カウンターは自動的に合います。quad_multsに、ヘッドの数は入りません。
位置の数だけに比例する取り分を数える
matvec(M, x, ctr)・per_position_mults(d_model, d_ff, ctr)・linear_mults(n, d_model, d_ff)を追加してください。per_position_multsは、線形写像6つ(射影4つ、フィードフォワード2つ)をそれぞれ実際に動かして乗算を数え、上げた分を返します。linear_multsは、その値に位置の数を掛けた整数です。
行列の値は何でもかまいません。数えることが目的なので、形だけ合わせた行列を作って動かせば十分です。6つは、d_model x d_modelが4つと、d_ff x d_modelが1つ、d_model x d_ffが1つです。フィードフォワードの2層目に渡すベクトルは、1層目が返した長さd_ffのものです。採点ツールは、カウンターが少なくとも6回以上呼ばれたかも見るので、式1つを一度に足して終わりにすると、不合格になります。linear_multsは、カウンターを受け取りません。
2つの取り分を並べる
cost_table(ns, d_model, d_ff)を作成してください。nsの長さごとに(n, quad, linear, quad + linear)の4つの欄を返し、順序はnsと同じです。
前に作ったquad_multsとlinear_multsをそのまま使えば、5行です。長さを2倍ずつ増やしながら読んでみてください。前のほうの欄は4倍ずつ、後ろのほうの欄は2倍ずつ増えます。最後の欄は、必ず前の2つの欄の合計である必要があります。片方だけ書いておくと、あとで交差点を探すときにずれます。
交差点を探す
crossover_n(d_model, d_ff)を作成してください。quad_mults(n, d_model) >= linear_mults(n, d_model, d_ff)が初めて真になるnです。等号を含みます。
1から上げながら探してもよく、手で解いてもかまいません。両辺をnと2 * d_modelで割ると、条件がとても短くなります。等号を抜かして不等号だけを使うと、答えがちょうど1つずれます。モデルの幅を変えながら呼び出してみてください。幅が広いほど、交差点が後ろへずれることが、数字で見えます。
マスクが捨てる半分を数える
causal_pairs(n)とwasted_pairs(n)を作成してください。causal_pairsは、因果マスクで実際に使われるスコアの数で、wasted_pairsは、n * nからその取り分を引いた値です。
i番目のクエリは、自分自身までしか見ません。行の長さが1、2、3と増えて、nで終わるので、数えてみると、ちょうどn * (n + 1) / 2です。対角線を除いてはいけません。自分自身は見るのが正しいのです。causal_pairs(0)は0、causal_pairs(1)は1です。捨てられる割合が、長さが長くなるほどどこへ近づくか、いくつか出力してみてください。
メモリと、1行ずつ増えるスコアを数える
score_bytes(n, n_heads, itemsize)・max_context_for_bytes(budget_bytes, n_heads, itemsize)・append_pairs(n)・generate_pairs(n, g)を作成してください。前の2つは、スコア行列をまるごと保持するときのバイト数と、バジェットに収まる最長のコンテキストで、後ろの2つは、1トークンを足すときに新しく生じるスコアの数と、g個を作る間の総和です。
score_bytesで、ヘッドの数を抜かしやすいので注意してください。max_context_for_bytesは、実数の平方根を四捨五入すると、1つ超えてしまうことがあるので、math.isqrtを使ってください。バジェットをヘッドとバイトで割ってから、整数の平方根をとれば済みます。append_pairs(n)は、新しいクエリ1つが自分自身を含めてn + 1個を見るので、n + 1です。generate_pairsは、n+1からn+gまでの合計で、generate_pairs(0, N)がcausal_pairs(N)と同じかどうかを、必ず確認してください。
数えた数字を記録に残す
上の関数を実際に動かして、/root/work/tf-cost/cost_report.jsonにはd_model・d_ff・n_heads・itemsize・table・quad_ratio・linear_ratio・crossover_n・quad_at_crossover・linear_at_crossover・causal_n・causal_pairs・wasted_pairs・wasted_fraction・score_bytes_at_causal_n・max_context_1gib・append_pairs_at_causal_n・generate_pairs_at_causal_n・identity_okを、/root/work/tf-cost/cost_report.mdには## 무엇을 세었나 ## 두 배로 늘리면 무엇이 네 배가 되나 ## 교차점은 어디인가 ## 인과 마스크가 버리는 절반 ## 메모리와 한 토큰씩 늘어나는 점수の5つの節で記録してください(韓国語の見出しは順に、「何を数えたか」「2倍にすると何が4倍になるのか」「交差点はどこか」「因果マスクが捨てる半分」「メモリと、1トークンずつ増えるスコア」という意味です)。
数字は手で書かず、自分のコードを動かして得た値で埋めてください。表の長さは[128, 256, 512, 1024, 2048, 4096]で、quad_ratio・linear_ratioは、最後の2行の間の倍率です。causal_nは2048、データ型は2バイト、バジェットは1073741824、生成するトークンの数は256です。identity_okは、generate_pairs(0, causal_n) == causal_pairs(causal_n)が真かどうかです。記録には、測っていない数字を書かないでください。実際のモデルの秒数やGBは、ここで測ったことがありません。