本番環境でのLLM最適化: 精度、注意、アーキテクチャ

本番環境で大規模言語モデル (LLMs) をデプロイするには、数十億のパラメータによる巨大なVRAM需要と、長い入力シーケンスに関連する二次的なメモリ増加という2つの主要なボトルネックを克服する必要があります。これに対処するため、Hugging Faceは低精度量子化、最適化された注意アルゴリズム、戦略的なアーキテクチャ選択の組み合わせを推奨します。

低精度によるメモリフットプリントの削減

数値精度を下げることで、モデルの重みをロードするために必要なVRAMが削減され、より小規模またはアクセスしやすいハードウェアでより大規模なモデルを実行できるようになります。

精度別のVRAM要件

短いテキスト入力(1024トークン未満)では、モデルの重みのロードがメモリコストの主な要因となります。VRAM要件の一般的な目安は次のとおりです:

  • float32: X億パラメータのモデルに対して、おおよそ4 * X GBのVRAMが必要です。
  • bfloat16/float16: X億パラメータのモデルに対して、おおよそ2 * X GBのVRAMが必要です。

たとえば、Llama-2-70bはbfloat16で約140GBのVRAMを必要とし、これは単一のA100(80GB)の容量を超え、テンソルまたはパイプライン並列処理が必要になります。

量子化(8ビットと4ビット)

量子化により精度を8ビットまたは4ビットにさらに低下させ、テキスト生成における精度への影響を最小限に抑えながらメモリ使用量を大幅に削減します。これは、テキスト生成が正確な値ではなく、次のトークンのロジットの相対分布に依存しているためです。

  • 8-bit 量子化: VRAM使用量を大幅に削減します(例:OctoCoderのピークメモリが約32GBから約15GBに低下)。計算中の動的デクォンタイズの必要性により、推論時にわずかな遅延が生じる可能性があります。
  • 4-bit 量子化: さらにVRAMを削減します(例:OctoCoderが約9.5GBに低下)。これにより、RTX 3090などのコンシューマーGPUでモデルを実行できるようになります。ただし、8ビット量子化よりも精度の顕著な劣化と推論の遅延が生じる可能性があります。

Flash Attentionによる推論の高速化

標準の自己注意は、シーケンス長 $N$ に対して計算とメモリの複雑さが二次的に増加し、長いコンテキスト(例:16,000トークン以上)では実行コストが非常に高くなります。

Flash Attentionアルゴリズム

Flash Attentionは、計算をより小さなチャンクに分割し、複数のソフトマックスステップを繰り返すことで注意機構を最適化します。これにより、大きな $QK^T$ 行列の作成を回避し、$N$ に対するメモリコストが二次的ではなく線形に増加するようになります。

パフォーマンスの向上

Flash Attentionは、ソフトマックス正規化統計の再計算によりFLOPsが増加しますが、遅い高帯域幅メモリ(VRAM)へのアクセスを最小限に抑え、高速なオンチップSRAMの使用を最大化するため、実際には高速になります。デフォルトの自己注意アルゴリズムと数値的に同じ出力を生成します。

長いコンテキストとチャットのためのアーキテクチャ最適化

トレーニング中のアーキテクチャの選択は、長いシーケンスやマルチターンダイアログをどれだけ効率的に処理するかを決定します。重要な2つの領域は、位置埋め込みとキー・バリュー(KV)キャッシュです。

相対位置埋め込み

絶対位置埋め込み(サインusoidalまたは学習済み)は、長いテキストでは性能が低く、トレーニング長を超えて外挿することが困難です。相対位置埋め込みはより効果的です:

  • Rotary Position Embedding (RoPE): クエリ-キー ペアを回転させることで位置をエンコードします。Falcon、Llama、PaLMで使用されています。
  • ALiBi: 事前に定義された値でスケールされた負の整数を $QK^T$ 行列に加算します。MPTとBLOOMで使用され、RoPEよりも長いシーケンスに対してより効果的に外挿することが一般的です。

キー・バリュー(KV)キャッシュの最適化

自己回帰生成では、KVキャッシュを使用してすべての以前のトークンのキー・バリューベクトルを保存し、毎ステップで再計算する必要をなくします。これにより、$QK^T$ の計算がベクトル-行列乗算 ($\text{query} \times \text{KV cache}$) に変換され、速度が大幅に向上します。

ただし、KVキャッシュはメモリのボトルネックになる可能性があります。このオーバーヘッドを削減する2つのアーキテクチャがあります:

  • Multi-Query Attention (MQA): すべての注意ヘッドで共有される単一のキー・バリュー投影ヘッドを使用します。これによりキャッシュサイズが劇的に削減されます(例:OctoCoderにおける16,000トークンシーケンスで15GBから400MB未満に減少)し、メモリ帯域幅のボトルネックも軽減されます。Falcon、PaLM、MPT、BLOOMで使用されています。
  • Grouped-Query Attention (GQA): MQAと標準のマルチヘッド注意の間の中間的なアプローチです。少数のKV投影ヘッド(例:2、4、または8)を使用することで、MQAよりもモデル容量を維持しつつ、その効率の大部分を保持します。Llama-2で使用されています。

Sources