Context Parallelism
コンテキスト並列化によって長いシーケンスを複数のGPUに分散し、メモリ使用量を削減してTransformerのトレーニングをスケールさせ、長文書や動画を扱うAIワークロードを支援する方法を学びます。
コンテキスト並列化は、長い入力シーケンスを複数のアクセラレーターに分割する分散コンピューティング手法です。各GPUはシーケンスの一部だけを処理し、アテンション中に他のGPUと連携します。これによりデバイスごとのアクティベーションメモリが削減され、非常に長い文書、長時間の動画、大量の画像パッチなど、1台のGPUのメモリを超える可能性がある入力でTransformerをトレーニングできます。
モデルのコンテキストウィンドウを単に大きくする方法とは異なり、コンテキスト並列化は、アーキテクチャが理論上処理できる情報量を変えません。代わりに、シーケンス次元を分散することで、そのコンテキストの処理を計算上実用的にします。
コンテキスト並列化の仕組み#
シーケンスに32,000トークンが含まれ、4台のGPUでコンテキスト並列化を行うとします。各デバイスには最初に約8,000トークンが割り当てられ、それに対応する中間アクティベーションが保存されます。
正規化やフィードフォワード層などの多くの演算は、各デバイス上のローカルなシーケンスチャンクを個別に処理できます。課題となるのはアテンション機構です。ローカルのクエリは、他のすべてのデバイスが保持するキーと値に対してアテンションを適用する必要がある場合があります。
そのため、実装ではキー・バリュー、つまりKVブロックをGPU間で交換します。リングアテンションでは、各デバイスがローカルデータで部分的なアテンションを計算し、KVブロックを次のデバイスに渡す処理を、シーケンス全体を処理するまで繰り返します。PyTorchのコンテキスト並列化チュートリアルでは、分散スケールドドットプロダクトアテンションを通じてこの動作を紹介しています。また、NVIDIAのコンテキスト並列化パッケージでは、all-gather、reduce-scatter、リングベースの通信方式について説明しています。
最終結果は、通常の数値的な差を除けばシーケンス全体のアテンションと数学的に等価ですが、すべてのシーケンスアクティベーションを1台のGPUに保持する必要はありません。
コンテキスト並列化が重要な理由#
長いシーケンスは、2つの大きなスケーリング上の問題を引き起こします。1つ目は、シーケンス長が伸びるにつれて保存されるアクティベーションがより多くのメモリを消費することです。2つ目は、標準的なセルフアテンションがシーケンス全体でトークンを比較するため、計算量と一時データが大きくなることです。
コンテキスト並列化は、アクティベーションをデバイス間で分割してメモリの問題に対処します。アテンションの処理も分割できますが、その分、通信コストが発生します。NCCLの集合通信操作ガイドに記載されているような演算を使い、効率的なシステムではKV転送と計算を並行して実行します。
この手法は、モデルの重みやバッチサイズではなく、シーケンス長が原因でメモリ不足エラーが発生する場合に最も有効です。分散トレーニングという幅広い分野に属し、複数の次元を同時にスケールさせるため、他の手法と組み合わせてよく使われます。
コンテキスト並列化と関連手法の比較#
- テンソル並列化は、個々の層内の演算や重み行列を分割します。一方、コンテキスト並列化はシーケンス方向にトークンを分割します。
- パイプライン並列化は、異なるモデル層のグループを異なるデバイスに割り当てます。分割するのはシーケンス長ではなく、モデルの深さです。
- データ並列化はモデルを複製し、それぞれのレプリカに異なるトレーニング例を割り当てます。PyTorchの分散処理の概要では、モデル全体と各サンプルが1台のGPUに収まる場合に、この手法を推奨しています。
- シーケンス並列化では、テンソル並列化に関連する一部の演算で、アクティベーションを分割することがよくあります。コンテキスト並列化では、ネットワークの入力とアクティベーション全体にわたって、より広くシーケンスを分割します。
これらの手法は相互に補完できます。NVIDIAの並列化戦略ガイドでは、コンテキスト、テンソル、パイプライン、データの各並列化を組み合わせて、多次元のデバイス配置を構成する方法を示しています。
実際のアプリケーション#
-
長文書を扱うAI: 法務や医療の言語モデルでは、案件ファイル、患者の病歴、技術マニュアル全体を処理する必要がある場合があります。コンテキスト並列化では、数千に及ぶ文書トークンを複数のアクセラレーターに分散し、離れたセクション間のアテンションを維持しながら、アクティベーションメモリの負荷を抑えます。
-
長時間の動画とマルチモーダル理解: 動画Transformerや大規模ビジョンモデルでは、フレーム、画像パッチ、音声セグメント、テキストを1つの長いトークンシーケンスとして表現する場合があります。そのシーケンスを分散すると、フレーム数や空間的な詳細を大幅に削減せずに、長時間の録画をモデルで解析できます。AWS Neuronのコンテキスト並列化の概要では、このような長いコンテキストのワークロードに向けて、アクセラレーターグループがKVシャードを交換する方法を紹介しています。
Ultralytics YOLO26のようなコンパクトなコンピュータービジョンアーキテクチャでは、通常、コンテキスト並列化は不要です。Ultralyticsのモデル学習ワークフローを通じた標準的なマルチGPUデータ並列化のほうが、一般にトレーニングの高速化に適しています。
実践的な使い方とトレードオフ#
次の最小限の例では、PyTorchの実験的なコンテキスト並列化APIとスケールドドットプロダクトアテンションの制御機能を使用します。cp_example.pyという名前で保存し、torchrun --standalone --nproc-per-node=2 cp_example.pyを使って2台のGPUで起動します。
import os
import torch
import torch.distributed as dist
import torch.nn.functional as F
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.experimental import context_parallel
from torch.nn.attention import SDPBackend, sdpa_kernel
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(rank)
torch.cuda.manual_seed(0)
dist.init_process_group("nccl")
mesh = init_device_mesh("cuda", (world_size,))
qkv = [torch.randn(1, 4, 4096, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True) for _ in range(3)]
with sdpa_kernel(SDPBackend.FLASH_ATTENTION), context_parallel(mesh, buffers=tuple(qkv), buffer_seq_dims=(2, 2, 2)):
output = F.scaled_dot_product_attention(*qkv, is_causal=True)
output.float().square().mean().backward()
dist.destroy_process_group()ここでは、次元2がシーケンス次元であるため、各プロセスはシーケンスシャードを受け取り、アテンションはデバイスメッシュ全体で連携して処理されます。
実際には、通信のオーバーヘッドを上回るメモリ削減効果があることを確認し、高速なインターコネクトを使用して、代表的なシーケンス長でベンチマークを行う必要があります。標準的なビジョンプロジェクトでは、Ultralytics Platformがデータセットのアノテーション、トレーニング、デプロイ、モニタリングのためのよりシンプルなクラウドおよびローカルのワークフローを提供するため、コンテキスト並列化を手動で設定する必要はありません。









