メインコンテンツまでスキップ

05. 長系列の分散学習(Context Parallelism)

04までは1台の GPU の中の話でした。しかし 100 万トークンの系列 は、1台の GPU メモリには到底収まりません。このページでは、長い系列を 複数の GPU に分散 して学習する コンテキスト並列(context parallelism, CP) と、StripedHyena 2 の畳み込みをどう分散するか——とりわけ美しい FFT 畳み込み ——を見ていきます。

1. なぜ分散学習が必要か

コンテキスト並列(CP)は、増大するモデルと入力次元に対応するため、系列を分割して複数デバイスで処理 する技術です。data / tensor / sequence / pipeline parallelism といった他の分散手法を補完します。

具体的には、入力 [D,L][D, L](次元 DD ×系列長 LL)を NcpN_{cp} 台のデバイスに 系列次元で分割 し、各デバイスが [D,L/Ncp][D,\, L/N_{cp}] のシャードを持ちます。

問題は、Self-Attention や畳み込みのような 系列混合(sequence mixing) が「他のデバイスにある部分」を必要とすることです。これをどう通信でやりくりするかが CP の核心です。

2. 2つの通信戦略:All-to-all と Point-to-point

論文は2つの通信戦略を整理します。

  • All-to-all(a2a) — 各デバイスが 全デバイスとデータを交換 し、系列全体を再構成します。[D,L/Ncp][D, L/N_{cp}] を再配置して [H/Ncp,L][H/N_{cp}, L](チャネルを分割し、系列は全長)にする。これで各ランクが系列混合を独立に実行できます。DeepSpeed Ulysses で使われる方式です。
  • Point-to-point(p2p) — 全デバイスにブロードキャストせず、1度に1つのピアと直接交換 します。ブロック計算と通信を何ラウンドも繰り返す。Self-Attention 版が有名な ring attention です。
All-to-all(全交換)D0D1D2D3Point-to-point(リング)D0D1D2D3
2つの通信戦略。a2a は全デバイス間で一括交換、p2p は隣接ピアとリング状に逐次交換する

ring attention:key/value をリングで回す

p2p の代表例 ring attention の動きを、もう少し細かく見てみましょう。各ランクは query・key・value をそれぞれ [D,L/Ncp][D,\, L/N_{cp}] で保持します。

  1. CP ランクを リング状 に並べる。各ランクは自分の query を持ったまま、key/value のかたまりを隣のランクへ順番に渡して いく。
  2. 各ステップで、手元の query と「いま受け取った key/value」で部分的な attention を計算する。
  3. softmax の統計を オンラインで逐次更新(online softmax)しながら積み上げる。
  4. NcpN_{cp} ステップ後、各ランクの query は すべての key/value を見終え、最終的な attention 結果の shard を保持する。

こうして 系列全体を1台に集めることなく full attention を完結 できます。a2a が「一括で全部集めてから計算」なのに対し、ring(p2p)は「少しずつ回しながら計算」するアプローチです。

3. Hyena 演算子をどう分散するか

Hyena の系列混合は、長さの異なる畳み込みです。CP 実装には演算子ごとの工夫が要ります。

  • a2a 畳み込み — 入力を a2a で [H/Ncp,L][H/N_{cp}, L] に再配置し、各シャードを CP 領域内で畳み込み、再び a2a で戻します。フィルタは各 CP 領域内に保存・生成します(SE は各ランクが H/NcpH/N_{cp} 個のフィルタを保持、MR/LI はフィルタ計算を領域内で実行)。逆伝播では勾配を元のランクへ戻すため、追加の a2a が2回必要です。
  • p2p 畳み込み — ここで FIR フィルタの局所性 が効きます。因果畳み込みの大部分は 通信なしでローカルに計算でき、各シャードのうち 最初の h1\ell_h - 1 要素だけ が「前のランクの末尾 h1\ell_h - 1 要素」を必要とします。各ランクはフィルタのコピーを持ち、全 DD チャネルの畳み込みを担当します。
ランク0:ローカル計算ランク1:ローカル計算ランク2:ローカル計算境界の ℓₕ−1 要素だけ前ランクから受信(赤)
p2p 畳み込みの局所性。FIR なので大半はローカル計算で済み、通信は境界の h1\ell_h-1 要素だけ

論文はさらに、通信と計算をオーバーラップ する拡張も提案します(境界以外をゼロ埋めで先に計算し、通信完了後に境界分を足し込む。これは 04 の2段階分解と同じ発想です)。

a2a 畳み込みについても、モデルとシーケンスが大きいほど 通信レイテンシがボトルネック になります(1回のメッセージが巨大なため)。論文はこれを緩和する channel-pipelined a2a も提案しています。入力を NpipeN_{pipe} 個のセグメントに分割し、CUDA ストリームで非同期に a2a を回しながら計算とオーバーラップ させることで、通信待ちの時間を計算で隠します。系列長方向ではなく チャネル方向にパイプライン する点が特徴です。

4. FFT 畳み込み:周波数領域の魔法

長い暗黙フィルタをもつ Hyena-LI は、FFT で計算するのが一般的です。鍵は 畳み込み定理——時間領域の畳み込みは、周波数領域では 単なる要素積 になる、という事実です。

xh=F1 ⁣(F(x)F(h))x * h = \mathcal{F}^{-1}\!\bigl(\mathcal{F}(x) \odot \mathcal{F}(h)\bigr)

ここで F,F1\mathcal{F}, \mathcal{F}^{-1} はフーリエ変換と逆変換です。離散フーリエ変換(DFT)は

y(k)=j=0l1x(j)ωljk,ωl=e2πi/ly(k) = \sum_{j=0}^{l-1} x(j)\,\omega_l^{jk}, \qquad \omega_l = e^{-2\pi i / l}

で定義され、素朴に計算すると O(l2)O(l^2) ですが、高速フーリエ変換(FFT)は分割統治で O(llogl)O(l \log l) を達成します。ll 点 DFT を2つの l/2l/2 点 DFT に分解する、という再帰がその正体です。

この「2分割」の演算が バタフライ(butterfly) です。Decimation-in-Frequency(DiF)バタフライは次の2演算からなります。

X=x+y,Y=(xy)ωjX = x + y, \qquad Y = (x - y)\,\omega^j
xy++×ωʲX = x+yY =(x−y)ωʲ
DiF バタフライ。2入力を足し引きし、片方にひねり係数 ωj\omega^j を掛ける。この「蝶」の積み重ねが FFT

なぜこれが分散に向くのか

FFT は「入力の独立な2分割に対して個別に FFT を行い、要素ごとの演算で結合する」構造です。これは 各分割を別デバイスに置く p2p の状況とぴったり一致 します。素朴にやると FFT 後の出力が ビット反転(bit-reversal)順 に並んでシャード配置が崩れますが、畳み込みでは FFT → 逆 FFT と往復する ので、DiF FFT と DiF 逆 FFT を組み合わせれば 入力と同じシャード配置に戻り、追加の a2a 通信が不要 になります。

Ncp>2N_{cp} > 2 への拡張には Radix-NN FFTll 点を NN 個の l/Nl/N 点に分解)を使い、N=NcpN = N_{cp} として分散 FFT 畳み込みを実現します。ただし論文は、さらなる最適化なしでは Hyena-LI には a2a の方が速いことが多かった とも正直に報告しています。

FFT の完全な導出は付録へ

DFT の定義から FFT の分割統治、バタフライ(DiF/DiT)、Radix-NN、そして分散 p2p FFT が追加通信なしで成立する理由まで——FFT 畳み込みの数学的な全体像は、付録 06. FFT 畳み込みの数学と分散実装 に独立してまとめました。数式でじっくり追いたい方はそちらへ。

LLM とのつながり:長文脈学習の共通基盤

ここで出てくる ring attention(p2p)DeepSpeed Ulysses(a2a) は、自然言語 LLM の長文脈学習でも標準的な技術です。StripedHyena 2 は、これらの考え方を 畳み込み演算子向けに拡張 し、さらに FFT という信号処理の道具を分散学習に持ち込んだ点が新しいと言えます。

5. 因果モデルの負荷分散

自己回帰(causal)モデルでは、計算が三角構造になり、単純に系列を分割すると デバイス間で負荷が偏ります(後ろのトークンほど計算量が多い)。これを避けるため、系列を工夫して割り当てます。

  • striped ordering(Brandon et al.):CP ランク数の2倍のシャードに分け、各ランクに2枚を縞状に配置。
  • zig-zag splitting(Llama 3, Dubey et al.):たとえば 8 シャード・Ncp=4N_{cp}=4 なら [x0,x7],[x1,x6],[x2,x5],[x3,x4][x_0, x_7], [x_1, x_6], [x_2, x_5], [x_3, x_4] のように、前半と後半をペアにして負荷を均等化。

StripedHyena 2 の学習では、Llama 3 と同じ zig-zag 分割を採用 しています。

6. まとめ:StripedHyena 2 の全体像

このセクションでは、Evo 2 を支える StripedHyena 2 を、論文に沿って深掘りしてきました。

ページ要点
01 概要hybridization-aware + hardware-aware な設計で、全域で速いマルチハイブリッドを実現
02 演算子Hyena 構造(射影+畳み込み+ゲーティング)と、SE/MR/LI の役割分担
03 設計SE-MR-LI レイアウト・filter grouping・1M 文脈拡張・1.2〜2.9 倍高速
04 カーネル畳み込み=Toeplitz、2段階分解で GEMM 化しテンソルコアをフル活用
05 分散学習(本ページ)a2a/p2p のコンテキスト並列、FFT 畳み込み、zig-zag 負荷分散

一貫しているのは、「演算子の設計」と「ハードウェア・分散アルゴリズム」を一体で考える(co-design) という思想です。Attention をただ別の演算子に置き換えるのではなく、短・中・長の畳み込みと attention を役割分担させ、それぞれを GPU で最速に走らせる ——この積み重ねが、Evo 240B・9 兆トークン・100 万文脈 という規模を可能にしました。

このセクションは今後も拡張予定

「LLM アーキテクチャ」は、StripedHyena 2 を第1弾として、Transformer を補完・代替する効率的アーキ(Mamba などの状態空間モデル、各種ハイブリッド)を今後追加していく予定です。

StripedHyena 2 が実際に何を成し遂げたかは、応用である Evo 2 で確かめられます。アーキテクチャの土台(Evo 2 アーキテクチャ編)と合わせて読むと、設計から応用までが一本につながります。

最後に原論文を:Ku, Nguyen, Romero et al., "Systems and Algorithms for Convolutional Multi-Hybrid Language Models at Scale", arXiv:2503.01868 (2025). arxiv.org/abs/2503.01868