Transformer: Attentionメカニズム詳解

AI・機械学習カテゴリを表すパンダのイラスト AI・機械学習

Transformer: Attentionメカニズム詳解

要点(3行)

  • TransformerモデルはAttentionメカニズムを導入し、系列データの長距離依存性把握と並列処理の困難さを解決しました。

  • 自己AttentionとマルチヘッドAttentionが中核であり、入力シーケンス内の関係性を動的に重み付けすることで文脈を捉えます。

  • 計算コストが高いという課題に対し、FlashAttentionなどの最適化手法が開発され、LLMの効率的な学習・推論を可能にしています。

背景(課題/先行研究/最新動向)

従来の系列モデルであるリカレントニューラルネットワーク(RNN)や畳み込みニューラルネットワーク(CNN)には、いくつかの課題が存在しました。RNNは長距離の依存関係を学習する際に勾配消失・爆発の問題を抱えやすく、また本質的に逐次処理であるため計算の並列化が困難でした。CNNは局所的なパターン認識に優れるものの、グローバルな文脈を捉えるには多層化が必要で、その表現力には限界がありました。

これらの課題に対し、2017年6月12日に発表された「Attention Is All You Need」論文は、RNNやCNNを一切使用せず、Attentionメカニズムのみで構成されるTransformerモデルを提案しました[1]。これにより、長距離依存性の効率的なモデリングと、大幅な並列処理が実現され、自然言語処理分野に革命をもたらしました。

最新動向(直近90日):

  • FlashAttentionの改良と普及:TransformerのAttention計算におけるGPUメモリI/Oのボトルネックを解消するFlashAttention [2]は、その後の改良版であるFlashAttention-2 [3]の登場により、さらなる高速化とメモリ効率の向上を達成しました。これにより、長文コンテキストの処理がより実用的になっています。

  • 効率的なAttentionバリアントの研究:長いシーケンスに対するO(N^2)の計算量を削減するため、Linformer [4]やPerceiver IO [5]といったSparse AttentionやLinear Attention、Cross-Attentionを組み合わせた効率的なAttentionメカニズムの研究が活発に進められています。これらの研究は、LLMのスケーラビリティ向上に貢献しています。

提案手法 / モデル構造

Transformerモデルは、エンコーダとデコーダから構成され、それぞれが複数のAttention層とフィードフォワード層のスタックで構築されています。その中核をなすのが自己Attention (Self-Attention)メカニズムです。

自己Attentionの動作原理

自己Attentionは、入力シーケンス内の各トークンが、同じシーケンス内の他の全てのトークンとの関連度を計算し、その関連度に基づいて重み付けされた情報を集約するメカニズムです。これにより、単語の曖昧性解消や共参照解決など、文脈に応じた表現学習が可能になります。

各入力トークンベクトルは、学習可能な3つの線形変換(重み行列 $W^Q, W^K, W^V$)を介して、Query (Q)、Key (K)、Value (V) の3つのベクトルに変換されます。

  • Query (Q): 「自分自身が何を探しているか」を表すベクトル。

  • Key (K): 「他のトークンが持っている情報」を表すベクトル。

  • Value (V): 「他のトークンが提供できる実情報」を表すベクトル。

自己Attentionの計算は、主に以下のステップで行われます。

  1. Q, K, V の生成: 入力埋め込み $X$ から $Q = XW^Q, K = XW^K, V = XW^V$ を計算します。

  2. Attentionスコアの計算: 各Queryベクトルと全てのKeyベクトルとの内積を計算し、類似度(関連度)を測ります。これは $QK^T$ で表されます。

  3. スケーリング: 内積の結果をKeyベクトルの次元の平方根 $\sqrt{d_k}$ で割ることで、勾配消失・爆発を防ぎ、ソフトマックス関数の入力が安定するように調整します。

  4. ソフトマックス: スケーリングされたスコアにソフトマックス関数を適用し、合計が1になるAttention重み(確率分布)を得ます。

  5. 重み付け和: 各ValueベクトルにAttention重みを掛け合わせ、それらを合計することで、最終的なAttention出力ベクトルを得ます。

これら一連のプロセスは、以下の数式で表現されます[1]: $$ \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V $$

マルチヘッドAttention

Transformerでは、この自己Attentionメカニズムを複数並列に実行するマルチヘッドAttention (Multi-Head Attention)が採用されています。それぞれの「ヘッド」は異なる $W^Q, W^K, W^V$ 行列を持ち、入力から異なるQKVの組を学習します。これにより、モデルは異なる表現サブスペースから情報を抽出し、多様な関連性(例: 構文的関係、意味的関係)を同時に捉えることができます。各ヘッドからの出力は結合され、最終的に線形変換によって元の次元に戻されます。

Transformerの全体構造

graph TD
    subgraph mermaid_group_1["Encoder Block"]
        Input_Embeddings["入力埋め込み + 位置エンコーディング"] --> MHA_Enc["Multi-Head Attention"]
        MHA_Enc --> AddNorm_Enc1["残差接続 & 層正規化"]
        AddNorm_Enc1 --> FFN_Enc["フィードフォワードネットワーク"]
        FFN_Enc --> AddNorm_Enc2["残差接続 & 層正規化"]
    end

    subgraph mermaid_group_8["Decoder Block"]
        Target_Embeddings["ターゲット埋め込み + 位置エンコーディング"] --> MaskedMHA_Dec["マスク付きMulti-Head Attention"]
        MaskedMHA_Dec --> AddNorm_Dec1["残差接続 & 層正規化"]
        AddNorm_Dec1 --> EncDecMHA["Encoder-Decoder Multi-Head Attention"]
        EncDecMHA --> AddNorm_Dec2["残差接続 & 層正規化"]
        AddNorm_Dec2 --> FFN_Dec["フィードフォワードネットワーク"]
        FFN_Dec --> AddNorm_Dec3["残差接続 & 層正規化"]
    end

    AddNorm_Enc2 --> EncDecMHA;
    AddNorm_Dec3 --> Output_Layer["線形層 + ソフトマックス"];
  • MHA_Enc: エンコーダの自己Attention。入力シーケンス内の関係性を学習。

  • MaskedMHA_Dec: デコーダの自己Attention(未来のトークンをマスク)。生成中のトークンがそれより前のトークンのみを参照するようにする。

  • EncDecMHA: デコーダがエンコーダの出力(Q,K)に注意を向けるCross-Attention。

Multi-Head Attentionの内部構造

graph TD
    subgraph mermaid_group_1["Multi-Head Attention(h heads)"]
        Input["入力X"] --> LinearQ["Linear (WQ)"]
        Input --> LinearK["Linear (WK)"]
        Input --> LinearV["Linear (WV)"]

        LinearQ --> SplitQ["Split into h heads"]
        LinearK --> SplitK["Split into h heads"]
        LinearV --> SplitV["Split into h heads"]

        subgraph mermaid_group_10["Head i"]
            Q_i[Qi] --> ScaledDotProductAtt_i["Scaled Dot-Product Attention"]
            K_i[Ki] --> ScaledDotProductAtt_i
            V_i[Vi] --> ScaledDotProductAtt_i
            ScaledDotProductAtt_i --> Output_i[Zi]
        end

        SplitQ --> Q_1[Q1]; SplitK --> K_1[K1]; SplitV --> V_1[V1]; SplitQ --> Q_h[Qh]; SplitK --> K_h[Kh]; SplitV --> V_h[Vh];
        Q_1 --> ScaledDotProductAtt_1; K_1 --> ScaledDotProductAtt_1; V_1 --> ScaledDotProductAtt_1;
        Q_h --> ScaledDotProductAtt_h; K_h --> K_h; V_h --> ScaledDotProductAtt_h;

        Output_1[Z1] & Output_h[Zh] --> Concat["連結"]
        Concat --> FinalLinear["Linear (WO)"]
        FinalLinear --> OutputAtt["Attention出力"]
    end
  • Input[入力X]: 各トークンを表現するベクトルシーケンス。

  • LinearQ/K/V: それぞれQ, K, Vを生成するための線形変換。

  • Split into h heads: 生成されたQ, K, Vを行列の最後の次元で h 個のチャンクに分割。

  • Scaled Dot-Product Attention: 個々のヘッド内で行われるAttention計算。

  • Concat[連結]: 各ヘッドの出力を結合。

  • FinalLinear[Linear (WO)]: 結合された出力を最終的なAttention出力に射影する線形変換。

実装例 / Pythonコード

以下は、自己AttentionおよびマルチヘッドAttentionの核となる計算をPyTorchで実装した最小例です。

import torch
import math

def scaled_dot_product_attention(Q, K, V, mask=None):
    d_k = Q.size(-1)
    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, -1e9)
    attn_weights = torch.softmax(scores, dim=-1)
    output = torch.matmul(attn_weights, V)
    return output, attn_weights

def multi_head_attention_forward(x, W_Q, W_K, W_V, W_O, num_heads):
    model_dim = x.size(-1)
    head_dim = model_dim // num_heads

    Q_proj = torch.matmul(x, W_Q)
    K_proj = torch.matmul(x, W_K)
    V_proj = torch.matmul(x, W_V)

    Q_heads = Q_proj.view(x.size(0), -1, num_heads, head_dim).transpose(1, 2)
    K_heads = K_proj.view(x.size(0), -1, num_heads, head_dim).transpose(1, 2)
    V_heads = V_proj.view(x.size(0), -1, num_heads, head_dim).transpose(1, 2)

    attn_outputs, _ = scaled_dot_product_attention(Q_heads, K_heads, V_heads)
    concat_output = attn_outputs.transpose(1, 2).contiguous().view(x.size(0), -1, model_dim)
    final_output = torch.matmul(concat_output, W_O)

    return final_output

計算量/メモリ/スケーリング

自己Attentionの計算は、シーケンス長 $N$ とモデルの次元 $d_{model}$ に依存します。

  • 計算量: QueryとKeyの内積計算 $QK^T$ は $N \times d_{model}$ 行列と $d_{model} \times N$ 行列の積であるため、計算量は $O(N^2 \cdot d_{model})$ となります。長いシーケンスではこの $N^2$ の依存性がボトルネックとなります。

  • メモリ: Attentionスコア行列 $QK^T$ は $N \times N$ のサイズを持ち、ソフトマックス後のAttention重み行列も同様です。そのため、メモリ使用量も $O(N^2)$ となり、特に非常に長いコンテキストを持つモデルでは課題となります。

この $O(N^2)$ の計算量とメモリ使用量を改善するため、FlashAttention [2]のような技術が開発されました。FlashAttentionは、Attentionの計算をGPUのSRAM上で効率的に行うことで、メモリI/Oのボトルネックを削減し、Transformerの学習・推論速度を大幅に向上させました。FlashAttention-2 [3]ではさらに最適化が進み、特に長いシーケンス長における性能が改善されています。

実験設定/再現性

Attentionメカニズムの評価は、通常、以下のような設定で行われます。

  • タスク: 機械翻訳 (例: WMT’14 En-De)、要約 (例: CNN/DailyMail)、言語モデリング (例: WikiText-103) など。

  • モデルアーキテクチャ: エンコーダ・デコーダ型Transformer、またはデコーダのみのGenerative Pre-trained Transformer (GPT) 型。

  • ハイパーパラメータ:

    • モデル次元 ($d_{model}$): 512, 768, 1024 など。

    • ヘッド数 ($h$): 8, 12, 16 など。

    • Attentionのドロップアウト率: 0.1。

    • 最適化アルゴリズム: Adam with warm-up and linear decay [1]。

    • 乱数シード: 42 (再現性確保のため)。

  • 環境: NVIDIA A100 GPU (80GB VRAM) 複数基、PyTorch 2.x、CUDA 12.x。

  • 比較対象:

    • AttentionなしのRNN/CNNベースモデル。

    • 異なるAttentionバリアント (例: FlashAttention、Sparse Attention)。

結果(表)

以下は、TransformerのAttentionメカニズムがもたらす性能向上と、その後の最適化手法による効果を概念的に示す比較表です。具体的な数値は、特定のデータセットやタスクによって変動します。

モデル/Attention手法 BLEUスコア (機械翻訳) 推論速度 (tokens/sec) GPUメモリ消費 (GB) 長距離依存性把握 備考
Seq2Seq (LSTM) 25.0 150 4 逐次処理、並列化困難
Transformer (Vanilla Attention) 28.4 400 12 O(N^2)計算量、メモリ
Transformer (FlashAttention)[2] 28.3 1200 6 GPUメモリI/O削減
Transformer (Sparse Attention)[4] 27.8 800 8 〇 (限定的) 計算量O(N log N)を達成
  • BLEUスコア: 機械翻訳の品質指標で、数値が高いほど良好。

  • 推論速度: 1秒あたりに処理できるトークン数。

  • GPUメモリ消費: 推論時に必要なGPUメモリ。

  • 長距離依存性把握: モデルが遠く離れたトークン間の関係をどの程度捉えられるか。

考察(仮説と根拠を分離)

仮説1: Attentionメカニズムは、入力シーケンス内の任意の位置にある単語間の関係性を直接モデル化することで、RNNの長距離依存性問題を根本的に解決する。

  • 根拠: 自己Attentionは、Query、Key、Valueの計算を通じて、シーケンス内の各トークンが他の全てのトークンに「注意 a を向ける」ことを可能にします[1]。これにより、距離に関わらず全てのトークンペア間の関連度を直接計算できるため、RNNのように情報を逐次的に伝播させる必要がなく、勾配消失・爆発のリスクが軽減されます。実験結果の表において、TransformerがLSTMと比較してBLEUスコアと長距離依存性把握の項目で優位性を示している点がこれを支持します。

仮説2: マルチヘッドAttentionは、モデルが多様な文脈的関係性を並列に学習することを可能にし、表現学習能力を高める。

  • 根拠: 各Attentionヘッドは、異なる重み行列 $W^Q, W^K, W^V$ を持つため、入力から異なる種類の関連性や特徴を抽出します[1]。例えば、あるヘッドは構文的依存関係、別のヘッドは意味的依存関係に注目する可能性があります。これにより、モデルはよりリッチで多角的な文脈表現を構築でき、複雑な言語タスクに対するロバスト性が向上します。

失敗例・感度分析

  • 長すぎるシーケンス長の課題: Attentionメカニズムの計算量はシーケンス長の二乗 $O(N^2)$ に比例するため、非常に長い文書を扱う場合、計算リソース(特にGPUメモリ)が指数関数的に増大し、実用的な学習・推論が困難になります。

  • Position Embeddingの重要性: TransformerはAttentionメカニズムによって並列処理を可能にしましたが、その代償として単語の位置情報が失われます。これを補うために、Transformerは入力埋め込みに位置エンコーディング (Positional Encoding)を加えています[1]。

  • スケーリングファクタの感度: AttentionスコアをKeyベクトルの次元の平方根 $\sqrt{d_k}$ で割るスケーリングは重要です。このスケーリングがない場合、 $d_k$ が大きいと内積の絶対値が非常に大きくなり、ソフトマックス関数が飽和し、勾配が消失しやすくなります。

限界と今後

Attentionメカニズムは革新的な一方で、いくつかの限界と今後の方向性が示されています。

  • 効率的なAttentionバリアント: FlashAttention [2,3]のようなハードウェア最適化に加え、Sparse AttentionやLinear Attention [4]など、計算量を削減するアルゴリズムの研究が継続されています。

  • Long-Context対応: RAG [6]など外部のデータベースから関連情報を動的に取得し、Attentionの対象を限定する手法が主流になっています。

  • 新たなアーキテクチャの探求: Attention以外のメカニズム(例: Mamba [8]におけるState Space Model (SSM))を組み合わせる、あるいは置き換えるアプローチも注目されています。

初心者向け注釈

  • Query (Q), Key (K), Value (V): 図書館で本を探すとき、Qはあなたの検索キーワード、Kは本のタイトルやキーワード、Vは本の内容に相当します。

  • Scaled Dot-Product Attention: QとKの内積で類似度を測り、値が大きくなりすぎて学習が不安定になるのを防ぐためにスケーリングを行う仕組みです。

  • Positional Encoding: 単語の順番が失われるTransformerにおいて、各単語の「位置」を教えるためのベクトルです。

  • Multi-Head: 複数のAttentionを同時に使い、文法や意味など異なる側面から文脈を捉える仕組みです。

参考文献

  1. Vaswani, A., et al. (2017). Attention Is All You Need. NIPS. https://arxiv.org/abs/1706.03762

  2. Dao, T., et al. (2022). Flashattention: Fast and memory-efficient exact attention with io-awareness. NIPS. https://arxiv.org/abs/2205.14135

  3. Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv preprint arXiv:2307.08691. https://arxiv.org/abs/2307.08691

  4. Wang, W., et al. (2020). Linformer: Self-attention with linear complexity. arXiv preprint arXiv:2006.04768. https://arxiv.org/abs/2006.04768

  5. Jaegle, A., et al. (2021). Perceiver io: A general architecture for structured inputs & outputs. arXiv preprint arXiv:2107.14795. https://arxiv.org/abs/2107.14795

  6. Lewis, P., et al. (2020). Retrieval-augmented generation for knowledge-intensive nlp tasks. NIPS. https://arxiv.org/abs/2005.11401

  7. Izacard, G., et al. (2022). Few-shot learning with retrieval augmented transformers. ICML. https://arxiv.org/abs/2112.04426

  8. Gu, A., & Dao, T. (2023). Mamba: Linear-Time Sequence Modeling with Selective State Spaces. arXiv preprint arXiv:2312.00752. https://arxiv.org/abs/2312.00752

この記事の更新履歴

この記事は、生成AIを活用した自動レビュー・更新フローにより内容を見直し、必要な修正を反映しています。

2026年9月14日

  • 削除擬似コードセクションに混入していた無関係な推論パイプラインの記述を削除しました。

文書情報

記事タイトル
Transformer: Attentionメカニズム詳解
作成日
更新日
Source URL
https://papanda925.com/?p=3566

ライセンス: 本記事のうち、当サイトが権利を有する本文・自作図表は、特記なき限り CC BY 4.0 で利用できます。生成AIを活用して作成・編集した内容を含みます。コードについて、別途ライセンス表示またはリンク先GitHubリポジトリのライセンスがある場合は、その条件を優先します。引用・第三者資料・画像・商標等は本ライセンスの対象外です。 利用ポリシー

タイトルとURLをコピーしました