Grouped Query Attention (GQA)
Grouped Query Attention(GQA)がKV-cacheのメモリ使用量を削減し、推論効率を高め、Transformerモデルの性能バランスを取る仕組みを学びます。
グループ化クエリアテンション (GQA) は、複数のクエリヘッドで、より少数のキーおよびバリューヘッドを共有するアテンション設計です。マルチヘッドクエリの多様な視点を維持しながら、キーとバリューに必要なメモリおよびデータ移動を削減します。GQAは、モデルが一度に1トークンずつ生成し、保存されたアテンション状態を繰り返し読み取る自己回帰推論で特に有効です。
GQAの仕組み#
各アテンションメカニズム内で、各トークンは3つの表現に射影されます。
- クエリ: 現在のトークンが求めている情報を表します。
- キー: 利用可能な各トークンが表す内容を示します。
- バリュー: クエリがキーと一致したときに取得される情報を含みます。
従来のセルフアテンションでは、モデルが異なる関係を並列に学習できるよう、これらの射影をヘッドに分割します。GQAでは多数のクエリヘッドを維持しつつ、それらをグループ化して共有のキーおよびバリューヘッドに割り当てます。たとえば、アテンションレイヤーで8個のクエリヘッドと2個のキーバリューヘッドを使用する場合があります。この場合、各キーバリューヘッドが4個のクエリヘッドを担当します。
アテンションの計算自体は、引き続きスケールドドット積アテンションです。変わるのはヘッドの構成と、生成、保存、読み取りを行うキーバリューデータの量です。そのため、GQAを有効にしたPyTorchのスケールドドット積アテンションなどのフレームワークでは、クエリヘッド数をキーバリューヘッド数で割り切れるようにする必要があります。(docs.pytorch.org)
GQAとマルチヘッドアテンションおよびマルチクエリアテンションの比較#
GQAは、マルチヘッドアテンションとマルチクエリアテンションの中間に位置します。
- マルチヘッドアテンション: 各クエリヘッドが独自のキーヘッドとバリューヘッドを持ちます。ヘッドの独立性は最大になりますが、キーバリューメモリの要件も最大になります。
- グループ化クエリアテンション: 複数のクエリヘッドが各キーバリューヘッドを共有します。表現能力とメモリ効率のバランスが取れています。
- マルチクエリアテンション: すべてのクエリヘッドが1つのキーヘッドと1つのバリューヘッドを共有します。キーバリューのストレージを最小限に抑えられますが、キーバリューの多様性は低くなります。
あるレイヤーに32個のクエリヘッドがある場合、マルチヘッドアテンションでも32個のキーバリューヘッドを使用する一方、GQAでは8個、マルチクエリアテンションでは1個を使用することがあります。GQAは、より広範なTransformerアーキテクチャの置き換えではなく、Transformerのアテンションレイヤー内で選択できる構成の1つです。NVIDIAのマルチヘッド、マルチクエリー、グループ化クエリアテンションに関するドキュメントでは、これらのバリエーションは、クエリヘッドを担当するキーバリューヘッドの数が主な違いであると説明されています。(nvidia.github.io)
GQAが重要な理由#
自己回帰生成では、モデルは以前のキーとバリューをKVキャッシュに保存します。標準的なマルチヘッドアテンションでは、すべてのアテンションヘッドが個別のキーとバリューをキャッシュに追加します。そのため、長いシーケンス、大きなバッチ、多数のモデルレイヤーによって、GPUメモリを大量に消費する可能性があります。
GQAではキーバリューヘッドの数が少ないため、キャッシュもヘッド数の削減率に比例して小さくなります。これにより、次の点を改善できる可能性があります。
- メモリ容量: より長いプロンプトや、より多くの同時リクエストを利用可能なメモリに収められます。
- 生成スループット: 生成される各トークンで読み取るキーバリューデータが少なくなります。
- 推論レイテンシ: メモリトラフィックの削減により、トークン間の処理時間を短縮できます。
- コンテキストウィンドウのスケーラビリティ: 長いシーケンスをより実用的に提供できるようになります。
これらの効果は、ハードウェア、シーケンス長、バッチサイズ、カーネルのサポート状況によって異なります。プロダクションランタイムでは、効率的なキャッシュ割り当ても必要です。たとえばNVIDIAのTensorRT-LLM KVキャッシュシステムは、GQAのサポートに加えて、ブロックベースのキャッシュ、再利用、オフロードを組み合わせています。(nvidia.github.io)
実世界での利用例#
長いコンテキストを扱うアシスタント: コーディングアシスタントは、回答を生成する前に数千個のソースコードトークンを処理する場合があります。GQAは、その履歴に関連付けられたキャッシュ済みのキーバリュー状態を削減するため、GPUメモリを比例して増やすことなく、サーバーでより長いファイルやより多くの同時ユーザーを処理できるようにします。
マルチモーダル画像・動画アシスタント: ビジョン言語モデルは、画像パッチや動画フレームを長いトークンシーケンスとして表現できます。GQAは、これらのビジュアルトークンが言語デコーダーに入力された後のアテンションキャッシュへの負荷を軽減し、追加の画像、フレーム、またはリクエストに使用できるメモリを増やせる可能性があります。これは、標準的なVision Transformerとは異なります。標準的なVision Transformerでは、自己回帰KVキャッシュを使用せず、すべての画像パッチを並列に処理する場合があります。
実装とトレードオフ#
PyTorchでは、enable_gqa=Trueを通じてGQAを利用できます。このドキュメントに記載されたCUDAの例では、8個のクエリヘッドと2個の共有キーバリューヘッドを使用します。
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
assert torch.cuda.is_available(), "A CUDA device is required."
query = torch.randn(1, 8, 16, 32, device="cuda")
key = torch.randn(1, 2, 16, 32, device="cuda")
value = torch.randn(1, 2, 16, 32, device="cuda")
# Four query heads share each key-value head.
with sdpa_kernel(SDPBackend.MATH):
output = F.scaled_dot_product_attention(
query,
key,
value,
enable_gqa=True,
)
print(output.shape)出力では8個のクエリヘッドが維持されますが、キーとバリューには2個のヘッドのみが使用されます。PyTorchアテンションバックエンドセレクターでは、処理を実行するサポート対象カーネルを制御できます。また、PyTorch Transformerビルディングブロックガイドでは、より幅広い実装コンテキストが提供されています。(docs.pytorch.org)
GQAは、モデルの設計時または適応時に選択する必要があり、既存のチェックポイントに対して推論時に一般的に切り替えられるオプションではありません。開発者は、フレームワークのサポート、ヘッド数の割り切れ方、数値的な挙動、メモリ使用量、タスク品質を検証する必要があります。コンピュータビジョンのデプロイでは、Ultralytics TensorRT統合やUltralyticsベンチマークモードなどの補完的なワークフローにより、アーキテクチャおよびランタイムの最適化が対象ハードウェア上で有意な改善をもたらすかどうかを測定できます。









