PISA、ブロック疎注意を O(N log N) の対数線形複雑度へ
長文脈で残るブロック選択の問題
長い文脈を扱う言語モデルでは、自己注意の計算量が大きな制約になります。標準的な注意機構は、すべてのクエリとキーの組み合わせを評価するため、系列長が増えると計算量は二次的に増加します。ブロック疎注意は重要なキーのブロックだけを残すことでこの負担を減らしますが、残すブロックを決める際に全ブロックを採点すれば、ルーティング自体が依然として二次コストになり得ます。
ByteDance Seed の研究チームが提案する PISA は、この選択段階を効率化するための手法です。中心にあるのは、すべての候補から一度に Top-K を選ぶのではなく、粗い表現を使って候補を段階的に絞り込むピラミッド型の探索です。
粗い表現から細かいブロックへ
PISA では、プーリングによってキーを複数の粒度で表現します。粗い階層では、より広い領域を少数の表現でまとめます。下位の階層へ進むにつれて、候補はより細かいブロックへ展開されます。この構造は系列長 N に対して O(log N) 個の階層から構成されます。
各クエリの探索は最も粗い階層から始まります。現在の階層では、全キーを対象にするのではなく、上限のある候補集合に対して LogSumExp に基づくスコアを計算し、上位の候補だけを次の階層へ送ります。この手続きを最細粒度まで繰り返すことで、最終的に疎注意で参照するブロックを決定します。
この設計では、各段階の探索規模が制御され、全ブロックを毎回走査する必要がなくなります。論文では、これにより手法全体の計算量を O(N log N) にできると説明しています。重要なのは、注意計算を疎にするだけでなく、疎化のための候補検索も階層化している点です。
実装と評価
PISA はアルゴリズムだけでなく、実行系も考慮しています。著者らは学習と推論の双方に対応する Triton カーネルを開発し、階層的なルーティングと LogSumExp のスコア計算を融合しました。また、クエリとキーの完全なスコア行列をメモリ上に materialize しない構成を採用しています。これにより、中間テンソルが必要とするメモリを抑えやすくなります。
評価は言語モデリングタスクを中心に行われています。概要によれば、常識推論などのベンチマークでは基準手法と同程度の性能を保ち、検索タスクではより良い結果が得られました。長い文脈から関連情報を探すタスクで改善が報告された点は、PISA の設計意図と整合します。ただし、提示された素材にはモデル設定、文脈長、絶対スコア、実際のスループットは含まれていないため、改善幅や適用範囲を断定することはできません。
意義と今後の論点
PISA は、ブロック疎注意における「どのブロックを選ぶか」を、注意計算とは別の重要な最適化対象として扱っています。長文脈モデルの効率化では、最終的な注意演算だけでなく、候補を探す段階でも密な計算を避ける必要があります。
一方、粗い表現に基づく早期判断にはリスクもあります。重要な情報を初期段階で候補から外すと、後の細粒度処理で回復できない可能性があります。したがって、候補数、プーリング方法、タスクごとの精度と計算量のバランスが今後の検証点になります。PISA は、長文脈注意を全面的な照合から階層的な検索へ移行する一つの設計例だと言えるでしょう。
コメント
ログイン状態を確認中…
コメントを読み込み中…