未来を形作るテクノロジーの深掘り記事。

Self-Attentionの先へ:Transformer後の世界

Transformerの二乗コストは深刻な課題。線形Attention・Attention residual・ハイブリッド構成が示す次世代アーキテクチャを解説。

絡まった糸がほどけ、整然とした光の流れへと変わっていく様子

Transformerアーキテクチャは、もう8年にわたってAIを動かし続けています。主要な言語モデルのほぼすべて、多くの画像生成システム、そして増え続ける音声・動画モデルが、「Attention Is All You Need」論文で発表された自己注意機構(self-attention)の上に成り立っています。ただ、self-attentionには根本的な問題があります。計算量とメモリ使用量がシーケンス長の二乗に比例するのです。入力長を2倍にすると、コストは4倍になります。

短いシーケンスなら問題になりません。しかし、私たちが目指している128Kトークンのコンテキストウィンドウ、そして人々が求めている100万トークンのウィンドウになると、深刻なボトルネックになります。その解決に向けて、さまざまな研究が進んでいます。レイヤーをまたいで計算を再利用するAttention residual、二乗コストを解消する線形Attentionの亜種、Attentionと軽量な仕組みを組み合わせたハイブリッドアーキテクチャなどです。Transformerがなくなるわけではありませんが、確実に作り直されつつあります。

なぜSelf-Attentionは高コストなのか

代替案を理解するには、self-attentionが実際に何を計算しているのかを押さえる必要があります。N個のトークンからなるシーケンスがあるとき、self-attentionはトークンのあらゆる組み合わせについて関連度スコアを計算します。トークン1とトークン2、トークン1とトークン3、……トークン1とトークンN、次にトークン2と他のすべてのトークン、という具合です。つまり、N²組のペアを扱うことになります。

import torch
import torch.nn.functional as F
def self_attention(Q, K, V):
"""
Standard self-attention.
Q, K, V: (batch, seq_len, d_model)
The attention matrix is seq_len × seq_len.
For seq_len = 1024:   ~1M entries   (manageable)
For seq_len = 32768:  ~1B entries   (expensive)
For seq_len = 131072: ~17B entries  (very expensive)
"""
d_k = Q.size(-1)
# This matmul creates the N×N attention matrix
scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
weights = F.softmax(scores, dim=-1)
return torch.matmul(weights, V)

4Kトークンなら、Attention行列は1,600万要素。最新のGPUなら何の問題もありません。128Kトークンになると、160億要素になります。100万トークンでは1兆を超えます。Flash Attention(計算量自体は減らないが、メモリアクセスのパターンを劇的に改善する手法)を使っても、最終的には二乗のスケーリングが勝ってしまいます。

そのため初期のTransformerは512や1024トークンが上限でした。ハードウェアと最適化の進歩によって上限は引き上げられてきましたが、私たちは数学的な壁と戦い続けています。線形スケーリング(O(N))が二乗スケーリング(O(N²))よりも根本的に優れていることは明らかで、多くの代替アーキテクチャがそこを目指しています。

Attention Residual:計算済みの結果を使い回す

Attentionのコストを下げる最も実用的なアプローチの一つは、Attentionを置き換えるのではなく、前のレイヤーの計算を再利用して各Attentionレイヤーを軽くする方法です。

ある観察があります。32層のような深いTransformerでは、隣接するレイヤーのAttentionパターンが驚くほど似ていることが多いのです。たとえば15層目と16層目は、わずかな調整を除けばほぼ同じ位置に注目しがちです。毎回N²のAttention行列をゼロから計算するのは冗長で、その作業の多くは1層前ですでに済んでいるのです。

Attention residualはこの性質を利用し、「このレイヤーが注目したい対象」と「前のレイヤーが計算した結果」の差分、すなわち「残差(residual)」のAttentionパターンを計算します。差分が小さければ(中間層では通常小さくなります)、計算は安く済みます。完全なAttentionパターンは、前のレイヤーのパターンに現在のレイヤーの残差を足し合わせることで得られます。

これは動画圧縮の仕組みと似ています。各フレームを独立して保存するのではなく、キーフレームを1つ保存し、そこからの差分(残差)を並べて保存するのです。差分はフレーム全体よりずっと小さいので、圧縮効率が劇的に上がります。

実際、Attention residualは深いモデルの中間層において、品質への影響をほとんど出さずに、Attentionの計算コストを30〜50%削減できると報告されています。最初と最後の数層は、パターンがより特徴的なため、依然として完全なAttention計算が必要です。しかし、大多数を占める中間層では大幅な高速化が得られます。

線形Attention:二乗コストからの脱却

線形Attentionの亜種は、Attentionを再定式化してO(N²)ではなくO(N)でスケールさせようとするものです。基本的な考え方は、N×NのAttention行列を明示的に計算するのをやめ、線形演算で同じ(または近似的に同じ)出力を得る方法を探すことです。

数学的なトリックは、softmaxのカーネル分解に依存しています。標準的なAttentionはsoftmax(QK^T)Vを計算します。softmaxをφ(Q) · φ(K)^Tと分解できる別のカーネル関数で置き換えると、計算順序を入れ替えられます。(φ(Q) · φ(K)^T) · V(N×Nの中間結果が生じる)の代わりに、φ(Q) · (φ(K)^T · V)(中間結果はd×d。dはモデル次元)を計算すればよいのです。長いシーケンスではd << Nなので、これは劇的に安く済みます。

def linear_attention(Q, K, V, feature_map=None):
"""
Linear attention via kernel feature maps.
Cost: O(N * d^2) instead of O(N^2 * d)
"""
if feature_map is None:
# ELU+1 is a common choice (from Katharopoulos et al.)
feature_map = lambda x: F.elu(x) + 1
Q = feature_map(Q)  # (batch, seq_len, d)
K = feature_map(K)  # (batch, seq_len, d)
# Key insight: compute K^T @ V first (d × d matrix)
# instead of Q @ K^T first (N × N matrix)
KV = torch.einsum('bnd,bnm->bdm', K, V)  # (batch, d, d)
# Then multiply by Q
output = torch.einsum('bnd,bdm->bnm', Q, KV)  # (batch, N, d)
# Normalize
Z = torch.einsum('bnd,bd->bn', Q, K.sum(dim=1))  # normalization
output = output / Z.unsqueeze(-1)
return output

注意点があります。softmaxを別のカーネル関数に置き換えると注目の分布が変わるため、softmax Attentionで学習したモデルが線形Attentionにそのまま移せるとは限りません。品質差は大きく縮まりました。最新の線形Attention亜種はsoftmax Attentionの95〜98%の品質に到達しています。それでも差は残っており、特に長距離の正確な検索が必要なタスクでは顕著です。

状態空間モデル:別のパラダイム

Mambaのような状態空間モデル(SSM)は、根本的に異なるアプローチを取ります。トークン同士の組み合わせを計算する代わりに、再帰的な処理でシーケンスを順に読み込み、トークンごとに更新される固定サイズの隠れ状態を保持します。これは本質的にO(N)です。トークンが2倍になれば処理時間も2倍で済み、4倍にはなりません。

最近のSSMの革新は、再帰のパラメータを入力に依存させる点にあります(selective state spaces)。これによって、二乗コストを払うことなく、内容に基づくAttentionに近い性質が得られます。どの情報を覚えて、どれを忘れるかを「選択」できるのです。Mamba系のモデルは多くのベンチマークでTransformerと同等の品質を示しつつ、長いシーケンスではかなり高速です。

トレードオフもあります。SSMはトークンを逐次処理するため、全トークンを同時に処理できるTransformerに比べて学習時の並列化が難しいのです。学習効率は重要です。推論が2倍速くても、学習が3倍遅ければ、全体の計算量の大部分は学習に使われるので、必ずしも得とは言えません。

ハイブリッドアーキテクチャ:現実的な道

現在の本番モデルでは、異なるAttention機構を組み合わせたハイブリッドアーキテクチャが主流になりつつあります。理由はシンプルで、モデルの部位によって適した計算の種類が違うからです。

  • グローバルな推論には完全なAttention。 数千トークン離れた関連コンテキストを見つけるために、シーケンス全体にまたがって注目が必要な層があります。こうした層には、標準の(場合によってはFlashで最適化された)self-attentionを使います。
  • 近傍のコンテキストにはローカルAttention。 多くの層は主に近くのトークンに注目します(スライディングウィンドウAttention)。256〜1024トークンの固定ウィンドウを使えば、コストはウィンドウサイズWを使ってO(N·W)まで下がります。
  • 広い文脈の集約には線形Attention。 一部の層はシーケンス全体の情報を集約する必要がありますが、厳密なAttention重みまでは不要です。線形Attentionなら、O(N)のコストでこれを実現できます。
  • 逐次的な処理にはSSM層。 Mamba系の層は、Attentionの計算をまったく使わずに、逐次的な依存関係を効率よく処理できます。

Jamba(AI21)や各種の研究用アーキテクチャは、層の役割に応じてこれらの機構を切り替えています。初期の層はローカルAttentionを使い(構文や局所的なパターンを処理)、中間層は線形AttentionやSSMを使い(より広い表現を構築)、戦略的に配置された数層だけが完全なAttentionを使います(大域的な推論と検索)。これにより、全体としてほぼ線形のスケーリングを実現しつつ、一部に完全なAttentionが必要なモデル品質も維持できます。

開発者が注目すべきポイント

言語モデルの上でアプリケーションを作っているなら、内部で起きているアーキテクチャの変化は、実際の開発に具体的な影響を与えます。

  • コンテキストウィンドウは広がり続けます。 Attentionのコストが下がれば、コンテキストウィンドウも拡大します。すると設計も変わります。関連情報を4Kウィンドウに収めるために複雑なRAGパイプラインを組む代わりに、100万トークンのプロンプトに全部詰め込めばよくなるかもしれません。シンプルさは魅力的ですが、レイテンシやコストへの影響はアーキテクチャによって異なります。
  • レイテンシの傾向が変わります。 Transformerのレイテンシは、ある地点までは比較的横ばいで、その後二乗的に増加します。線形AttentionやSSMのモデルは、より緩やかで線形的な増加を示します。応答時間が重要なアプリケーションでは、モデルのスケーリング特性を理解することが重要です。
  • 品質の差はタスク依存です。 線形Attentionモデルは、長いコンテキストの特定位置からの正確な検索(「47ページのリストの3番目の項目は何だったか?」のような質問)を要するタスクでは、わずかに劣る可能性があります。一方、一般的な理解を求めるタスクでは同等に機能します。ユースケースを把握しておきましょう。
  • 推論の最適化がより重要になります。 モデルが異なるAttention機構を混在させるほど複雑になると、推論エンジンは異種の計算を効率よく扱う必要があります。vLLMやTensorRT-LLMといったフレームワークも対応を進めていますが、独自アーキテクチャはすぐにはサポートされないかもしれません。

Transformerは置き換えられるのではなく、進化しているのです。Self-attentionは、トークン同士の関係をモデル化する手段として、今なお最も表現力の高い仕組みです。ただし、すべての層で、常にN²の完全コストで使う必要はありません。今後数年のモデルは、Attentionを外科手術のように使い分けるでしょう。最も重要な箇所には完全な精度を、それ以外にはより安価な代替手段を。その結果、現在の品質に並ぶか、それを超えつつ、より高速で、より長いコンテキストを扱え、運用コストも下がったモデルが生まれるはずです。注目に値する動きです。