Linear Attention / Performerの理論と実装 — カーネル近似で計算量をO(n)にする

Transformerは自然言語処理や画像認識で驚異的な性能を発揮していますが、ひとつ決定的な弱点があります。それはSelf-Attentionの計算量が系列長 $n$ の二乗に比例するという点です。系列長が512なら問題ありませんが、4096になると計算量は64倍、16384になれば1024倍に膨れ上がります。論文全文やゲノム配列、長時間音声のように数万トークン以上を一度に扱いたい場面では、この $O(n^2)$ のスケーリングが致命的なボトルネックになります。

この問題に対して、「Attentionの計算を厳密に $O(n)$ にできないか」という野心的なアプローチがLinear Attentionです。鍵となるアイデアは、Softmax Attentionをカーネル関数として再解釈し、行列積の結合順序を変えるだけで二次の計算量を線形に落とすというものです。2020年にKatharopoulos et al.が提案した線形Attention、そしてChoromanski et al.が提案したPerformer(FAVOR+)は、ランダム特徴量を使ってSoftmax Attentionをカーネル近似し、理論的にも実用的にも $O(n)$ のAttentionを実現しました。

Attentionの計算量 総当たりO(n^2)から要約経由O(n)への概念図

この記事の核心を1枚にまとめたのが上の図です。左の標準Attentionは全クエリ・全キーのペアを総当たりで結ぶため、計算が $n^2$ に比例します。右のLinear Attentionは、いったんキーとバリューを「要約表」$\bm{\Phi}_K^\top \bm{V}$ に集約し、各クエリはその要約表を引くだけにします。総当たりが「集計 $n$ 回+検索 $n$ 回」に置き換わり、計算量が $O(n)$ に落ちる — これがこの記事を通して理解するゴールです。

Linear Attentionを理解すると、以下のような応用・知見が得られます。

  • 超長系列の効率的処理: 数万〜数十万トークンの入力を、GPUメモリを気にせず処理できるアーキテクチャの設計原理がわかる
  • RNNとTransformerの統一的理解: Linear Attentionは再帰的に計算可能であり、TransformerとRNNの境界が実はなめらかにつながっていることがわかる
  • カーネル法の新しい応用: 機械学習の古典的手法であるカーネルトリックが、最新のTransformerアーキテクチャに自然に現れることで、理論的な視野が広がる
  • State Space ModelやRetentionへの橋渡し: Mamba、RetNetなど近年の線形計算量モデルは、Linear Attentionの考え方を発展させたものであり、ここで学ぶ概念がそのまま活きる

本記事の内容

  • 標準Attentionの $O(n^2)$ ボトルネックの再確認
  • カーネルトリックによるAttentionの再定式化
  • 行列の結合順序の変更による $O(n)$ 計算の実現
  • Performer: FAVOR+(Fast Attention Via Orthogonal Random features)の理論
  • Random Feature Map $\phi(\bm{x})$ と正ランダム特徴量
  • 近似誤差と精度のトレードオフ
  • Linear AttentionのRNN的解釈と再帰計算
  • Pythonでの標準Attention vs Linear Attentionの実装と計算時間の比較

前提知識

この記事を読む前に、以下の記事を読んでおくと理解が深まります。

画像なし
Self-Attentionの理論と実装
Query・Key・Valueの計算とScaled Dot-Product Attentionの導出
画像なし
Multi-Head Attentionの理論と実装
複数のAttentionヘッドで異なる部分空間の情報を統合する仕組み
画像なし
Sparse Attention(Longformer・BigBird)の理論と実装
スパースパターンによるAttentionの効率化手法
画像なし
Flash Attentionの理論と実装
GPUメモリ階層を活用したAttentionの高速化手法

標準Attentionの計算量問題

$O(n^2)$ のボトルネック

Linear Attentionの動機を理解するために、まず標準的なSelf-Attentionの計算コストを正確に振り返りましょう。

Scaled Dot-Product Attentionは次の式で定義されます。

$$ \begin{equation} \text{Attention}(\bm{Q}, \bm{K}, \bm{V}) = \text{softmax}\left(\frac{\bm{Q}\bm{K}^\top}{\sqrt{d_k}}\right)\bm{V} \end{equation} $$

ここで $\bm{Q}, \bm{K} \in \mathbb{R}^{n \times d_k}$、$\bm{V} \in \mathbb{R}^{n \times d_v}$ です。$n$ は系列長、$d_k$ はQuery/Keyの次元、$d_v$ はValueの次元です。

この計算の中で最もコストが高いのが、$\bm{Q}\bm{K}^\top$ の計算です。$n \times d_k$ の行列と $d_k \times n$ の行列の積なので、結果は $n \times n$ の行列になります。つまり、系列長のペア数に比例した $O(n^2 d_k)$ の計算量と $O(n^2)$ のメモリが必要です。

なぜ $O(n^2)$ が問題なのか

具体的な数字で考えてみましょう。$d_k = 64$(一般的なヘッド次元)として、注意行列 $\bm{A} \in \mathbb{R}^{n \times n}$ のメモリ使用量(FP32)を見ます。

系列長 $n$ 注意行列の要素数 メモリ量(FP32) 用途の例
512 262,144 1 MB BERT標準
4,096 16,777,216 64 MB GPT-2
16,384 268,435,456 1 GB 長文書処理
65,536 4,294,967,296 16 GB 論文全体
131,072 17,179,869,184 64 GB ゲノム配列

これは1ヘッド・1レイヤー分です。Multi-Head Attention(例えば12ヘッド×12レイヤー)全体では、この144倍のメモリが必要になります。系列長65,536ではGPU1枚に収まらないことは明らかです。

注意行列メモリが系列長の二乗で爆発する図

両対数でプロットすると、注意行列のメモリが系列長 $n$ に対してまっすぐ傾き2で増えていく様子がはっきり見えます。系列長512ではわずか1MBですが、論文全体に相当する65,536トークンでは16GBに達し、灰色の破線で示したGPU 40GBの壁をあっという間に超えてしまいます。この急増こそが、長系列を扱ううえで $O(n^2)$ が致命的になる理由です。

3つの効率化アプローチ

この $O(n^2)$ 問題に対して、大きく3つのアプローチが提案されてきました。

  1. Sparse Attention(Longformer、BigBird): 注意パターンをスパースにして、実質的な計算量を $O(n)$ にする。ただし、どのペアに注意するかのパターン設計が必要
  2. FlashAttention: 計算量自体は $O(n^2)$ のままだが、GPUメモリ階層を最適化して高速化する。近似なしの厳密計算
  3. Linear Attention: Attentionの定式化を変更し、数学的に $O(n)$ の計算量を実現する

本記事で扱うLinear Attentionは、3つ目のアプローチです。Sparse Attentionがパターンの工夫で計算を省略するのに対し、Linear Attentionは行列積の結合則という代数的な性質を利用して、同等(または近似的に同等)の計算を少ない計算量で実現します。

では、この「行列積の結合則」をAttentionに適用するために、まずSoftmax Attentionをカーネル関数の視点から再定式化しましょう。

カーネルトリックによるAttentionの再定式化

Attentionを1つのクエリから見る

行列表記のAttention式を、まず1つのクエリベクトルに対する出力として書き直してみましょう。$i$ 番目のクエリ $\bm{q}_i$ に対するAttention出力は次のとおりです。

$$ \begin{equation} \text{Attn}(\bm{q}_i) = \frac{\sum_{j=1}^{n} \exp\left(\bm{q}_i^\top \bm{k}_j / \sqrt{d_k}\right) \bm{v}_j}{\sum_{j=1}^{n} \exp\left(\bm{q}_i^\top \bm{k}_j / \sqrt{d_k}\right)} \end{equation} $$

分母はSoftmaxの正規化定数です。$\bm{v}_j$ はValue行列の $j$ 行目のベクトルです。

ここで重要な見方があります。$\exp(\bm{q}_i^\top \bm{k}_j / \sqrt{d_k})$ という部分は、$\bm{q}_i$ と $\bm{k}_j$ の類似度を非負の値で測る関数です。これはカーネル関数そのものです。

カーネル関数とは

カーネル法は機械学習の古典的な概念ですが、ここでのポイントは非常にシンプルです。カーネル関数 $\kappa(\bm{x}, \bm{y})$ とは、2つのベクトルの「類似度」を非負の実数値で返す関数です。代表的な例として、RBFカーネル $\kappa(\bm{x}, \bm{y}) = \exp(-\|\bm{x} – \bm{y}\|^2 / 2\sigma^2)$ があります。

多くのカーネル関数には、ある特徴写像 $\phi: \mathbb{R}^d \to \mathbb{R}^D$ が存在し、次のように分解できます。

$$ \begin{equation} \kappa(\bm{x}, \bm{y}) = \phi(\bm{x})^\top \phi(\bm{y}) \end{equation} $$

つまり、高次元空間 $\mathbb{R}^D$ に写像してから内積を取ることで、元の空間では表現できない複雑な類似度を計算できるのです。これがカーネルトリックの基本的な考え方です。

カーネルトリック 類似度を特徴写像phiの内積で表す概念図

図のように、元の入力空間にあるベクトル $\bm{q}, \bm{k}$ をいったん特徴写像 $\phi$ で高次元の特徴空間へ送り、そこで内積を取ったものがカーネル値 $\kappa(\bm{q}, \bm{k})$ になります。ポイントは、複雑な類似度の計算が「写像してから内積を取るだけ」という単純な操作に分解できる点です。後で見るように、この分解こそが行列積の結合順序を入れ替えて計算量を落とすための鍵になります。

Softmax Attentionのカーネル解釈

Attention式のSoftmaxをカーネル関数として捉え直しましょう。表記を簡単にするため、スケーリング $1/\sqrt{d_k}$ はQueryまたはKeyに吸収済みとします。

Softmax Attentionの $i$ 番目の出力は以下のように書けます。

$$ \begin{equation} \text{Attn}(\bm{q}_i) = \frac{\sum_{j=1}^{n} \kappa_{\text{sm}}(\bm{q}_i, \bm{k}_j) \bm{v}_j}{\sum_{j=1}^{n} \kappa_{\text{sm}}(\bm{q}_i, \bm{k}_j)} \end{equation} $$

ここで $\kappa_{\text{sm}}(\bm{q}, \bm{k}) = \exp(\bm{q}^\top \bm{k})$ です。

この形はNadaraya–Watson推定量(カーネル回帰)と全く同じ構造であることに気づきます。カーネル回帰では、クエリ点 $\bm{q}_i$ に近いデータ点(ここでは $\bm{k}_j$)に大きな重みを与えて、対応する値($\bm{v}_j$)の重み付き平均を取ります。

一般のカーネル関数による汎化

ここで視点を大きく広げます。Softmax Attentionの $\exp(\bm{q}^\top \bm{k})$ は、カーネル関数の1つの選択肢にすぎません。任意の正定値カーネル $\kappa$ を使って、一般化されたAttentionを定義できます。

$$ \begin{equation} \text{Attn}_\kappa(\bm{q}_i) = \frac{\sum_{j=1}^{n} \kappa(\bm{q}_i, \bm{k}_j) \bm{v}_j}{\sum_{j=1}^{n} \kappa(\bm{q}_i, \bm{k}_j)} \end{equation} $$

もしこのカーネル関数が有限次元の特徴写像 $\phi$ で分解できるなら、つまり $\kappa(\bm{q}, \bm{k}) = \phi(\bm{q})^\top \phi(\bm{k})$ と書けるなら、式を次のように変形できます。

分子を書き下すと、

$$ \sum_{j=1}^{n} \kappa(\bm{q}_i, \bm{k}_j) \bm{v}_j = \sum_{j=1}^{n} \phi(\bm{q}_i)^\top \phi(\bm{k}_j) \bm{v}_j $$

$\phi(\bm{q}_i)$ は $j$ に依存しないため、和の外に出せます。

$$ = \phi(\bm{q}_i)^\top \sum_{j=1}^{n} \phi(\bm{k}_j) \bm{v}_j^\top $$

ここで、$\phi(\bm{k}_j) \in \mathbb{R}^D$ と $\bm{v}_j \in \mathbb{R}^{d_v}$ の外積 $\phi(\bm{k}_j) \bm{v}_j^\top \in \mathbb{R}^{D \times d_v}$ を全ての $j$ について足し合わせた行列を用意します。この変形が決定的に重要なポイントです。

同様に、分母は $\phi(\bm{q}_i)^\top \sum_{j=1}^{n} \phi(\bm{k}_j)$ と書けます。

この変形がなぜ重要なのかを、次のセクションで具体的に見ていきましょう。計算量が劇的に変わります。

行列の結合順序の変更による $O(n)$ 計算

結合則のトリック — 核心的アイデア

Linear Attentionの核心は、実にシンプルな代数的事実に基づいています。行列積には結合則(associativity)が成り立つため、$(AB)C = A(BC)$ です。しかし、計算量は結合の順序によって大きく異なります

日常的な例で説明しましょう。3つの行列 $\bm{A} \in \mathbb{R}^{1000 \times 2}$、$\bm{B} \in \mathbb{R}^{2 \times 1000}$、$\bm{C} \in \mathbb{R}^{1000 \times 5}$ を掛けるとします。

  • $(\bm{A}\bm{B})\bm{C}$: まず $\bm{A}\bm{B} \in \mathbb{R}^{1000 \times 1000}$(計算量 $2 \times 10^6$)、次に $\times \bm{C}$(計算量 $5 \times 10^6$)。合計 $\approx 7 \times 10^6$ FLOPS
  • $\bm{A}(\bm{B}\bm{C})$: まず $\bm{B}\bm{C} \in \mathbb{R}^{2 \times 5}$(計算量 $10^4$)、次に $\bm{A} \times$(計算量 $10^4$)。合計 $\approx 2 \times 10^4$ FLOPS

結合順序を変えるだけで、計算量が350倍も違います。Linear Attentionは、まさにこの原理をAttentionに適用します。

同じ積でも結合順序で約350倍の演算回数差

棒グラフは対数スケールなので一見の差以上に開きがあります。$(AB)C$ のように先に $n \times n$ の大きな行列を作ってしまう順序では約 $7 \times 10^6$ 回の演算が必要ですが、$A(BC)$ のように先に小さな行列を作る順序ではわずか $2 \times 10^4$ 回で済みます。同じ答えを得るのに、巨大な中間行列を作るかどうかだけでここまで差がつくのです。

標準Attention: $(\bm{Q}\bm{K}^\top)\bm{V}$ の計算量

標準Attentionでは、まずスコア行列を計算してからValueを掛けます。

$$ \text{Output} = \underbrace{(\bm{Q}\bm{K}^\top)}_{n \times n} \bm{V} $$

  • $\bm{Q}\bm{K}^\top$: $(n \times d_k) \times (d_k \times n) \to n \times n$。計算量 $O(n^2 d_k)$
  • $(\bm{Q}\bm{K}^\top)\bm{V}$: $(n \times n) \times (n \times d_v) \to n \times d_v$。計算量 $O(n^2 d_v)$

合計: $O(n^2 d)$(ただし $d = d_k = d_v$ と仮定)

中間に $n \times n$ の行列が出現するため、メモリも $O(n^2)$ です。

Linear Attention: $\bm{Q}(\bm{K}^\top\bm{V})$ の計算量

カーネル分解 $\kappa(\bm{q}, \bm{k}) = \phi(\bm{q})^\top \phi(\bm{k})$ を使うと、Softmax正規化を含むAttention出力全体が次のように書けます。$\bm{\Phi}_Q, \bm{\Phi}_K \in \mathbb{R}^{n \times D}$ を、各行が $\phi(\bm{q}_i)^\top$、$\phi(\bm{k}_j)^\top$ である行列とします。

$$ \begin{equation} \text{Output}_i = \frac{\phi(\bm{q}_i)^\top \sum_{j=1}^{n} \phi(\bm{k}_j) \bm{v}_j^\top}{\phi(\bm{q}_i)^\top \sum_{j=1}^{n} \phi(\bm{k}_j)} \end{equation} $$

行列表記では、

$$ \text{Output} = \text{diag}^{-1}\!\left(\bm{\Phi}_Q \bm{\Phi}_K^\top \bm{1}_n\right) \cdot \bm{\Phi}_Q (\bm{\Phi}_K^\top \bm{V}) $$

ここで $\text{diag}^{-1}(\cdot)$ は正規化のための対角行列です。重要なのは括弧の位置です。

$\bm{\Phi}_K^\top \bm{V}$ を先に計算すると、

  • $\bm{\Phi}_K^\top \bm{V}$: $(D \times n) \times (n \times d_v) \to D \times d_v$。計算量 $O(nDd_v)$
  • $\bm{\Phi}_Q \times (\bm{\Phi}_K^\top \bm{V})$: $(n \times D) \times (D \times d_v) \to n \times d_v$。計算量 $O(nDd_v)$

合計: $O(nDd_v)$

特徴写像の次元 $D$ がモデル次元 $d$ と同程度であれば、計算量は $O(nd^2)$、つまり系列長 $n$ に対して線形です。中間に $n \times n$ の行列は一切出現しないため、メモリも $O(n)$ です。

結合順序 標準(QK^T)V vs Linear Q(K^TV) の行列形状と計算量

上下の図は、まったく同じ積を異なる順序で計算したときに現れる中間行列を比べたものです。上段の標準 $(\bm{Q}\bm{K}^\top)\bm{V}$ では、ピンクで示した $n \times n$ の巨大な注意行列が中間に現れ、これが $O(n^2)$ の計算量とメモリの正体です。下段のLinear $\bm{\Phi}_Q(\bm{\Phi}_K^\top \bm{V})$ では、中間に現れるのは黄色の $D \times d_v$ という系列長に依存しない小さな要約行列だけで、$n \times n$ 行列はどこにも現れません。形状を追うだけで、なぜ計算量が線形に落ちるのかが目で確認できます。

直感的な理解 — 要約してから検索する

この計算量の違いを直感的に理解するアナロジーを考えましょう。

標準Attentionは「全校生徒に個別に質問する」方法です。生徒が $n$ 人いて、$n$ 個の質問があれば、質問回数は $n^2$ 回です。

Linear Attentionは「まず全生徒の回答を集計して要約表を作り、各質問にはその要約表から答える」方法です。要約表の作成に $n$ 回、各質問への回答に $n$ 回、合計 $2n$ 回で済みます。

$\bm{\Phi}_K^\top \bm{V} \in \mathbb{R}^{D \times d_v}$ がまさにこの「要約表」に相当します。系列全体の情報を $D \times d_v$ の行列に圧縮してから、各クエリで検索するわけです。

ただし、ここにはひとつ大きな問題が残っています。Softmax Attentionのカーネル $\kappa_{\text{sm}}(\bm{q}, \bm{k}) = \exp(\bm{q}^\top \bm{k})$ に対して、有限次元の特徴写像 $\phi$ は厳密には存在しないのです。$\exp(\bm{q}^\top \bm{k})$ はテイラー展開すると無限次元の特徴空間を必要とします。では、どうすればよいのでしょうか? ここで登場するのがPerformerの FAVOR+ です。

Performer: FAVOR+ の理論

FAVOR+ の基本思想

Choromanski et al.(2021)が提案したPerformerは、Softmax Attentionのカーネル $\exp(\bm{q}^\top \bm{k})$ をランダム特徴量(random features)で有限次元に近似するという手法です。FAVOR+はFast Attention Via positive Orthogonal Random featuresの略です。

アイデアの出発点は、Rahimi & Recht(2007)のRandom Fourier Featuresです。任意のシフト不変カーネルはランダムな三角関数の内積で近似できるという定理ですが、FAVOR+はこれをAttentionに特化した形で発展させています。

ランダム特徴量による近似

まず、$\exp(\bm{q}^\top \bm{k})$ を近似可能な形に変形します。次の恒等式が出発点です。

$$ \exp(\bm{q}^\top \bm{k}) = \exp\!\left(\frac{\|\bm{q}\|^2}{2}\right) \cdot \exp\!\left(-\frac{\|\bm{q} – \bm{k}\|^2}{2}\right) \cdot \exp\!\left(\frac{\|\bm{k}\|^2}{2}\right) $$

この等式が成り立つ理由を確認しましょう。右辺の中央の指数部分を展開すると、

$$ -\frac{\|\bm{q} – \bm{k}\|^2}{2} = -\frac{\|\bm{q}\|^2 – 2\bm{q}^\top\bm{k} + \|\bm{k}\|^2}{2} $$

3つの指数関数の積の指数部分をまとめると、

$$ \frac{\|\bm{q}\|^2}{2} – \frac{\|\bm{q}\|^2}{2} + \bm{q}^\top\bm{k} – \frac{\|\bm{k}\|^2}{2} + \frac{\|\bm{k}\|^2}{2} = \bm{q}^\top\bm{k} $$

となり、確かに左辺と一致します。

中央の $\exp(-\|\bm{q} – \bm{k}\|^2 / 2)$ はガウスカーネル(RBFカーネル)です。ガウスカーネルのランダム特徴量近似は既知の結果として利用できます。

Random Feature Map の構成

ガウスカーネルのランダム特徴量近似では、$\bm{\omega}_1, \ldots, \bm{\omega}_m \sim \mathcal{N}(\bm{0}, \bm{I}_d)$ を独立にサンプリングし、次の近似を使います。

$$ \exp\!\left(-\frac{\|\bm{q} – \bm{k}\|^2}{2}\right) \approx \frac{1}{m} \sum_{l=1}^{m} \exp(i\bm{\omega}_l^\top \bm{q}) \exp(-i\bm{\omega}_l^\top \bm{k}) $$

しかし、この三角関数ベースの近似には問題があります。結果が負になりうるのです。Attentionの重みは非負でなければなりませんから、近似値が負になると注意重みの意味が破綻してしまいます。

Positive Random Features — FAVOR+ の鍵

FAVOR+は、この問題を「正ランダム特徴量(positive random features)」で解決します。三角関数の代わりに、指数関数を使った特徴写像を定義します。

$\bm{\omega}_1, \ldots, \bm{\omega}_m \sim \mathcal{N}(\bm{0}, \bm{I}_d)$ として、特徴写像 $\phi: \mathbb{R}^d \to \mathbb{R}^m_+$ を次のように定義します。

$$ \begin{equation} \phi(\bm{x}) = \frac{\exp\!\left(-\frac{\|\bm{x}\|^2}{2}\right)}{\sqrt{m}} \begin{pmatrix} \exp(\bm{\omega}_1^\top \bm{x}) \\ \exp(\bm{\omega}_2^\top \bm{x}) \\ \vdots \\ \exp(\bm{\omega}_m^\top \bm{x}) \end{pmatrix} \end{equation} $$

この特徴写像の各成分は $\exp$ の結果なので常に正です。そして、この特徴写像の内積がSoftmax Attentionのカーネルを近似することを確認しましょう。

$\phi(\bm{q})^\top \phi(\bm{k})$ を計算します。

$$ \phi(\bm{q})^\top \phi(\bm{k}) = \frac{\exp\!\left(-\frac{\|\bm{q}\|^2}{2}\right) \exp\!\left(-\frac{\|\bm{k}\|^2}{2}\right)}{m} \sum_{l=1}^{m} \exp(\bm{\omega}_l^\top \bm{q}) \exp(\bm{\omega}_l^\top \bm{k}) $$

各 $\bm{\omega}_l$ は独立にガウス分布からサンプリングされているため、$m \to \infty$ で大数の法則により、

$$ \frac{1}{m}\sum_{l=1}^{m} \exp(\bm{\omega}_l^\top \bm{q}) \exp(\bm{\omega}_l^\top \bm{k}) \to \mathbb{E}_{\bm{\omega} \sim \mathcal{N}(\bm{0}, \bm{I})}\left[\exp(\bm{\omega}^\top \bm{q}) \exp(\bm{\omega}^\top \bm{k})\right] $$

この期待値を閉じた形で計算します。$\bm{\omega} \sim \mathcal{N}(\bm{0}, \bm{I}_d)$ のとき、

$$ \mathbb{E}\left[\exp(\bm{\omega}^\top \bm{q}) \exp(\bm{\omega}^\top \bm{k})\right] = \mathbb{E}\left[\exp\!\left(\bm{\omega}^\top (\bm{q} + \bm{k})\right)\right] $$

ガウス分布のモーメント母関数 $\mathbb{E}[\exp(\bm{\omega}^\top \bm{t})] = \exp(\|\bm{t}\|^2 / 2)$ を使うと、

$$ = \exp\!\left(\frac{\|\bm{q} + \bm{k}\|^2}{2}\right) = \exp\!\left(\frac{\|\bm{q}\|^2 + 2\bm{q}^\top\bm{k} + \|\bm{k}\|^2}{2}\right) $$

したがって、

$$ \phi(\bm{q})^\top \phi(\bm{k}) \to \exp\!\left(-\frac{\|\bm{q}\|^2}{2}\right) \exp\!\left(\frac{\|\bm{q}\|^2 + 2\bm{q}^\top\bm{k} + \|\bm{k}\|^2}{2}\right) \exp\!\left(-\frac{\|\bm{k}\|^2}{2}\right) = \exp(\bm{q}^\top\bm{k}) $$

見事に $\exp(\bm{q}^\top \bm{k})$ に収束します。有限の $m$ では近似ですが、$m$ を大きくするほど精度が上がります。

FAVOR+ 正ランダム特徴量の非負性とカーネル値への収束

左の図は、正ランダム特徴量が必要な理由を示しています。三角関数ベースの特徴量(赤)は塗りつぶした領域のように負の値を取りうるため、注意重みの非負性が壊れてしまいます。一方、指数関数ベースの特徴量(緑)は常に正なので、Attentionの重みとして安心して使えます。右の図は、ランダム特徴量の本数 $m$ を増やすと推定値 $\phi(\bm{q})^\top\phi(\bm{k})$ が真のカーネル値 $\exp(\bm{q}^\top\bm{k})$(黒破線)へ収束していく様子で、$m$ が小さいうちは大きく振れますが、本数を増やすほどばらつきが収まっていきます。

直交ランダム特徴量(Orthogonal Random Features)

FAVOR+のもうひとつの工夫は、ランダムベクトル $\bm{\omega}_1, \ldots, \bm{\omega}_m$ を独立にサンプリングするのではなく、直交するように構成することです。

具体的には、$d \times d$ のランダム直交行列 $\bm{M}$(QR分解やハウスホルダー変換で生成)の行ベクトルに、$\chi^2$ 分布からサンプリングしたスケーリングを掛けてランダムベクトルを作ります。$m > d$ の場合は複数のランダム直交行列をブロック状に並べます。

直交ランダム特徴量の利点は、独立サンプリングより近似の分散が小さくなることです。直感的には、独立なサンプルでは似た方向のベクトルが偶然選ばれうるため空間のカバレッジにムラが生じますが、直交ベクトルは空間を均等にカバーするため、少ない本数 $m$ で効率よくカーネルを近似できます。

Choromanski et al.は、直交ランダム特徴量の推定量の分散が独立ランダム特徴量よりも厳密に小さくなることを理論的に証明しています。

直交ランダム特徴量 vs 独立サンプリングの空間カバレッジ

2次元で模式的に描くと違いがよくわかります。左の独立サンプリングでは、矢印(ランダムベクトルの方向)が偶然似た向きに固まったり、逆に空いた領域ができたりとムラが生じます。右の直交ランダム特徴量では、ベクトルが互いに直交するように配置されるため空間を均等にカバーでき、同じ本数 $m$ でもムラなくカーネルを近似できます。これが、少ない特徴量で低分散な近似を得られる直感的な理由です。

では次に、この近似の誤差がどの程度かを定量的に見ていきましょう。

近似誤差と精度のトレードオフ

近似誤差の理論的評価

FAVOR+の近似は $m$ 個のランダム特徴量で $\exp(\bm{q}^\top \bm{k})$ を推定するものなので、当然ながら有限の $m$ では誤差が生じます。その近似精度は、以下のように評価されます。

特徴量の次元 $m$ に対して、$\phi(\bm{q})^\top \phi(\bm{k})$ と $\exp(\bm{q}^\top \bm{k})$ の相対誤差は確率的に $O(1/\sqrt{m})$ の収束を示します。つまり、$m$ を4倍にすると誤差はおよそ半分になります。

精度と計算量のトレードオフ

特徴量の次元 $m$ は精度と計算量のバランスを決める重要なハイパーパラメータです。

  • $m$ が小さい: 計算量は少ないが近似誤差が大きく、Attentionパターンが歪む
  • $m$ が大きい: 近似精度は高いが、$O(nmd_v)$ の計算量が増加する

実用的には、$m$ はモデル次元 $d$ と同程度($m \approx d$ または $m \approx 2d$)に設定されることが多いです。Choromanski et al.の実験では、$m = d \log d$ 程度でSoftmax Attentionとほぼ同等の性能が得られると報告されています。

ただし、重要な注意点があります。近似誤差は $\bm{q}^\top \bm{k}$ の値が大きいペアほど増幅されます。$\exp(\bm{q}^\top \bm{k})$ は $\bm{q}^\top \bm{k}$ に対して指数関数的に増加するため、大きなスコアのペアの相対誤差は小さいものの、絶対誤差が大きくなりがちです。これは実際のタスクにおいて、特にAttentionがピーキー(少数のトークンに強く集中する)な場合に影響します。

Softmax Attentionとの性能比較

実践上のPerformerの位置づけをまとめると、

側面 標準Attention Performer (FAVOR+)
計算量 $O(n^2 d)$ $O(nmd)$
メモリ $O(n^2)$ $O(nm + md)$
近似誤差 なし(厳密) あり($m$ に依存)
因果マスク 自然にサポート 再帰計算で対応可能
短い系列 効率的 オーバーヘッドあり
長い系列 メモリ不足 効率的

短い系列($n \leq 1024$ 程度)では標準Attentionの方が高速な場合もありますが、系列長が大きくなるほどLinear Attentionの優位性が明確になります。

ここまでの議論は全系列を一括で処理する「バッチ」的な視点でした。しかし、Linear Attentionにはもうひとつ魅力的な性質があります — それはRNNのように再帰的に計算できることです。

Linear AttentionのRNN的解釈

再帰的な状態更新

標準のSoftmax Attentionは、全系列が揃わないと注意行列のSoftmax正規化を計算できないため、本質的に系列全体を一括処理する必要があります。しかし、Linear Attentionでは話が違います。

因果的(causal)な Linear Attention の出力を考えましょう。時刻 $t$ の出力は、過去の時刻 $1, \ldots, t$ のみに注目します。

$$ \begin{equation} \bm{y}_t = \frac{\phi(\bm{q}_t)^\top \sum_{j=1}^{t} \phi(\bm{k}_j) \bm{v}_j^\top}{\phi(\bm{q}_t)^\top \sum_{j=1}^{t} \phi(\bm{k}_j)} \end{equation} $$

ここで、$\bm{S}_t \in \mathbb{R}^{D \times d_v}$ と $\bm{z}_t \in \mathbb{R}^{D}$ を次のように定義します。

$$ \bm{S}_t = \sum_{j=1}^{t} \phi(\bm{k}_j) \bm{v}_j^\top, \quad \bm{z}_t = \sum_{j=1}^{t} \phi(\bm{k}_j) $$

すると、これらは次の再帰式で更新できます。

$$ \begin{align} \bm{S}_t &= \bm{S}_{t-1} + \phi(\bm{k}_t) \bm{v}_t^\top \\ \bm{z}_t &= \bm{z}_{t-1} + \phi(\bm{k}_t) \end{align} $$

出力は、

$$ \bm{y}_t = \frac{\phi(\bm{q}_t)^\top \bm{S}_t}{\phi(\bm{q}_t)^\top \bm{z}_t} $$

これは完全にRNNの構造です。$\bm{S}_t$ が「隠れ状態」に相当し、新しいトークンが来るたびに、そのKey-Valueペアの情報を加算的に蓄積していきます。

因果的Linear AttentionのRNN的解釈 隠れ状態S_tの逐次更新

図のように、各時刻で隠れ状態 $\bm{S}_t$ は前時刻の状態 $\bm{S}_{t-1}$ に、新しく来たトークンの外積 $\phi(\bm{k}_t)\bm{v}_t^\top$ を足し込むだけで更新されます。横方向の黒い矢印が状態の引き継ぎ、下からの青い矢印が新規入力、上への緑の矢印が出力 $\bm{y}_t$ です。Softmaxのように全系列を一度に見る必要がなく、トークンを1つずつ流すだけで計算が進む — これがLinear AttentionをRNNとして解釈できる根拠です。ただし加算のみで忘却がない点が、後述するRetentionやMambaの出発点になります。

RNNとの比較

この再帰構造を、従来のRNN(LSTMやGRU)と比較してみましょう。

従来のRNN: 隠れ状態 $\bm{h}_t \in \mathbb{R}^{d}$ は固定サイズのベクトルで、ゲート機構を通じて更新されます。情報の蓄積容量は $d$ 次元で固定です。

Linear AttentionのRNN: 隠れ状態 $\bm{S}_t \in \mathbb{R}^{D \times d_v}$ は行列です。$D \times d_v$ のサイズを持つため、従来のRNNよりもはるかに大きな情報容量を持てます。ただし、情報は加算的にしか蓄積されないため、忘却のメカニズムがありません。

忘却のないRNNの問題点

加算的な状態更新 $\bm{S}_t = \bm{S}_{t-1} + \phi(\bm{k}_t)\bm{v}_t^\top$ には「忘却」がありません。系列が長くなると、古いKey-Valueペアの情報が薄まらずに蓄積され続けるため、新しい情報が埋もれてしまう可能性があります。

この問題は近年の研究で活発に議論されており、減衰係数を導入したRetention(Sun et al., 2023)や、選択的なゲーティングを導入したMamba(Gu & Dao, 2023)など、Linear Attentionの考え方を発展させたモデルが次々と提案されています。

$$ \bm{S}_t = \gamma \bm{S}_{t-1} + \phi(\bm{k}_t) \bm{v}_t^\top \quad (\text{Retentionの場合、} 0 < \gamma < 1) $$

このように、Linear Attentionは理論的にRNNとTransformerを橋渡しする重要な位置にあり、効率的なシーケンスモデルの設計における基盤的なアイデアとなっています。

ここまでで理論を一通り見てきました。次は、実際にPythonで標準AttentionとLinear Attentionを実装し、動作と計算時間を比較してみましょう。

Pythonでの実装

標準Attentionの実装

まず、比較対象となる標準的なScaled Dot-Product Attentionを NumPy で実装します。

import numpy as np

def standard_attention(Q, K, V):
    """
    標準的なScaled Dot-Product Attention
    Q: (n, d_k), K: (n, d_k), V: (n, d_v)
    戻り値: (n, d_v)
    """
    d_k = Q.shape[1]
    # スコア行列の計算: (n, n)
    scores = Q @ K.T / np.sqrt(d_k)
    # 数値安定性のためmaxを引く
    scores = scores - np.max(scores, axis=1, keepdims=True)
    # Softmaxの計算
    exp_scores = np.exp(scores)
    attention_weights = exp_scores / np.sum(exp_scores, axis=1, keepdims=True)
    # Valueの重み付き和: (n, d_v)
    output = attention_weights @ V
    return output, attention_weights

ポイントは、Q @ K.T で $n \times n$ のスコア行列を明示的に構築している点です。このステップが $O(n^2)$ のメモリを消費します。

ELU特徴写像によるLinear Attentionの実装

最もシンプルなLinear Attentionとして、Katharopoulos et al.(2020)が提案したELU特徴写像を使った方法を実装します。この手法では $\phi(\bm{x}) = \text{elu}(\bm{x}) + 1$ という単純な特徴写像を使います。ELUは $\text{elu}(x) = x\ (x \geq 0)$, $e^x – 1\ (x < 0)$ ですので、$\text{elu}(x) + 1$ は常に正になります。

import numpy as np

def elu_feature_map(x):
    """ELU+1 特徴写像(常に正の値を返す)"""
    return np.where(x >= 0, x + 1, np.exp(x))

def linear_attention(Q, K, V):
    """
    Linear Attention(ELU特徴写像)
    Q: (n, d_k), K: (n, d_k), V: (n, d_v)
    戻り値: (n, d_v)
    """
    # 特徴写像の適用
    Q_prime = elu_feature_map(Q)  # (n, d_k)
    K_prime = elu_feature_map(K)  # (n, d_k)

    # Key-Valueの要約行列: K'^T V = (d_k, n)(n, d_v) -> (d_k, d_v)
    KV = K_prime.T @ V  # O(n * d_k * d_v)

    # 正規化定数: K'の列方向合計 -> (d_k,)
    Z = K_prime.sum(axis=0)  # O(n * d_k)

    # 各クエリに対する出力: Q' @ KV -> (n, d_v)
    numerator = Q_prime @ KV   # O(n * d_k * d_v)
    denominator = Q_prime @ Z  # O(n * d_k) -> (n,)

    # 正規化
    output = numerator / denominator[:, np.newaxis]
    return output

この実装で最も重要な行は KV = K_prime.T @ V です。$\bm{\Phi}_K^\top \bm{V}$ を先に計算することで、$n \times n$ の行列を一切作らずにAttention出力を得ています。

FAVOR+(Positive Random Features)の実装

次に、Performer の FAVOR+ を実装します。Softmax Attention を近似する正ランダム特徴量を使います。

import numpy as np

def favor_plus_feature_map(x, random_matrix):
    """
    FAVOR+ の正ランダム特徴写像
    x: (n, d_k) - 入力ベクトル
    random_matrix: (m, d_k) - ランダム射影行列
    戻り値: (n, m)
    """
    m = random_matrix.shape[0]
    # x を random_matrix で射影: (n, d_k) @ (d_k, m) -> (n, m)
    projection = x @ random_matrix.T  # (n, m)
    # xのノルムの二乗
    norm_sq = np.sum(x ** 2, axis=1, keepdims=True)  # (n, 1)
    # 正ランダム特徴量: exp(-||x||^2/2) * exp(w^T x) / sqrt(m)
    features = np.exp(projection - norm_sq / 2) / np.sqrt(m)
    return features

def create_orthogonal_random_matrix(m, d):
    """直交ランダム行列の生成"""
    # 必要なブロック数
    n_blocks = (m + d - 1) // d
    blocks = []
    for _ in range(n_blocks):
        # ガウス乱数行列を生成
        H = np.random.randn(d, d)
        # QR分解で直交行列を得る
        Q_mat, _ = np.linalg.qr(H)
        # chi分布のスケーリング
        S = np.sqrt(np.random.chisquare(d, size=d))
        blocks.append(S[:, np.newaxis] * Q_mat)
    # ブロックを結合してm行に切り詰め
    return np.vstack(blocks)[:m]

def performer_attention(Q, K, V, m=None, random_matrix=None):
    """
    Performer (FAVOR+) Attention
    Q: (n, d_k), K: (n, d_k), V: (n, d_v)
    m: ランダム特徴量の次元
    戻り値: (n, d_v)
    """
    d_k = Q.shape[1]
    if m is None:
        m = d_k  # デフォルトはd_kと同じ

    # スケーリングをQueryに吸収
    Q_scaled = Q / (d_k ** 0.25)
    K_scaled = K / (d_k ** 0.25)

    # 直交ランダム行列の生成
    if random_matrix is None:
        random_matrix = create_orthogonal_random_matrix(m, d_k)

    # 正ランダム特徴写像の適用
    Q_prime = favor_plus_feature_map(Q_scaled, random_matrix)  # (n, m)
    K_prime = favor_plus_feature_map(K_scaled, random_matrix)  # (n, m)

    # Linear Attention の計算(結合順序を変更)
    KV = K_prime.T @ V           # (m, d_v)
    Z = K_prime.sum(axis=0)      # (m,)
    numerator = Q_prime @ KV     # (n, d_v)
    denominator = Q_prime @ Z    # (n,)

    output = numerator / denominator[:, np.newaxis]
    return output

favor_plus_feature_map 関数の中で、np.exp(projection - norm_sq / 2) が正ランダム特徴量を計算しています。norm_sq / 2 を引くことで数値安定性を確保しつつ、前セクションで導出した $\exp(-\|\bm{x}\|^2/2) \cdot \exp(\bm{\omega}^\top \bm{x})$ を実装しています。

それでは、これらの実装が正しく動作するか、出力を比較してみましょう。

標準Attention と Linear Attention の出力比較

小さな例で近似精度を確認

まず、小さな系列長で3つの手法の出力を比較し、Linear Attention と FAVOR+ が標準Attentionの出力をどの程度近似しているかを確認します。

import numpy as np
import matplotlib.pyplot as plt

def standard_attention(Q, K, V):
    d_k = Q.shape[1]
    scores = Q @ K.T / np.sqrt(d_k)
    scores = scores - np.max(scores, axis=1, keepdims=True)
    exp_scores = np.exp(scores)
    attn_w = exp_scores / np.sum(exp_scores, axis=1, keepdims=True)
    return attn_w @ V, attn_w

def elu_feature_map(x):
    return np.where(x >= 0, x + 1, np.exp(x))

def linear_attention(Q, K, V):
    Q_prime = elu_feature_map(Q)
    K_prime = elu_feature_map(K)
    KV = K_prime.T @ V
    Z = K_prime.sum(axis=0)
    numerator = Q_prime @ KV
    denominator = Q_prime @ Z
    return numerator / denominator[:, np.newaxis]

def favor_plus_feature_map(x, W):
    m = W.shape[0]
    proj = x @ W.T
    norm_sq = np.sum(x ** 2, axis=1, keepdims=True)
    return np.exp(proj - norm_sq / 2) / np.sqrt(m)

def create_orthogonal_random_matrix(m, d):
    n_blocks = (m + d - 1) // d
    blocks = []
    for _ in range(n_blocks):
        H = np.random.randn(d, d)
        Q_mat, _ = np.linalg.qr(H)
        S = np.sqrt(np.random.chisquare(d, size=d))
        blocks.append(S[:, np.newaxis] * Q_mat)
    return np.vstack(blocks)[:m]

def performer_attention(Q, K, V, m=None):
    d_k = Q.shape[1]
    if m is None:
        m = d_k
    Q_s = Q / (d_k ** 0.25)
    K_s = K / (d_k ** 0.25)
    W = create_orthogonal_random_matrix(m, d_k)
    Q_prime = favor_plus_feature_map(Q_s, W)
    K_prime = favor_plus_feature_map(K_s, W)
    KV = K_prime.T @ V
    Z = K_prime.sum(axis=0)
    num = Q_prime @ KV
    den = Q_prime @ Z
    return num / den[:, np.newaxis]

# 再現性のためシードを固定
np.random.seed(42)

# テストデータ
n = 64   # 系列長
d_k = 32 # Key/Query次元
d_v = 32 # Value次元

Q = np.random.randn(n, d_k) * 0.5
K = np.random.randn(n, d_k) * 0.5
V = np.random.randn(n, d_v) * 0.5

# 3つの手法で出力を計算
out_std, attn_w = standard_attention(Q, K, V)
out_lin = linear_attention(Q, K, V)
out_perf = performer_attention(Q, K, V, m=64)

# 各位置の出力ベクトルのMSEを計算
mse_linear = np.mean((out_std - out_lin) ** 2, axis=1)
mse_performer = np.mean((out_std - out_perf) ** 2, axis=1)

fig, axes = plt.subplots(1, 3, figsize=(18, 5))

# 出力の1次元目を比較
axes[0].plot(out_std[:, 0], label='Standard Attention', linewidth=2)
axes[0].plot(out_perf[:, 0], '--', label='Performer (m=64)', linewidth=2)
axes[0].plot(out_lin[:, 0], ':', label='Linear (ELU)', linewidth=2)
axes[0].set_xlabel('Position')
axes[0].set_ylabel('Output (dim 0)')
axes[0].set_title('Output Comparison (1st dimension)')
axes[0].legend()
axes[0].grid(True, alpha=0.3)

# MSEの分布
axes[1].semilogy(mse_linear, 'o-', label='Linear (ELU)', alpha=0.7, markersize=3)
axes[1].semilogy(mse_performer, 's-', label='Performer (m=64)', alpha=0.7, markersize=3)
axes[1].set_xlabel('Position')
axes[1].set_ylabel('MSE (log scale)')
axes[1].set_title('Per-position MSE vs Standard Attention')
axes[1].legend()
axes[1].grid(True, alpha=0.3)

# 注意重み行列の可視化
im = axes[2].imshow(attn_w, cmap='viridis', aspect='auto')
axes[2].set_xlabel('Key position')
axes[2].set_ylabel('Query position')
axes[2].set_title('Standard Attention Weights')
plt.colorbar(im, ax=axes[2], shrink=0.8)

plt.tight_layout()
plt.savefig('linear_attention_comparison.png', dpi=150, bbox_inches='tight')
plt.show()

print(f"Linear (ELU) - Mean MSE: {np.mean(mse_linear):.6f}")
print(f"Performer (m=64) - Mean MSE: {np.mean(mse_performer):.6f}")
print(f"Standard Attention output range: [{out_std.min():.4f}, {out_std.max():.4f}]")

上のコードでは、系列長64、次元32の小さな例で3つの手法の出力を比較しています。左のグラフは出力ベクトルの1次元目を位置ごとにプロットしたもので、PerformerとELU Linear Attentionが標準Attentionの出力をどの程度再現しているかを視覚的に示します。中央のグラフは位置ごとのMSE(対数スケール)で、近似誤差の大きさを定量的に確認できます。右のグラフは標準Attentionの注意重み行列です。

Performerの出力は標準Attentionの出力に比較的近い値を示す一方、ELU特徴写像によるLinear Attentionは異なるカーネル関数を使っているため出力パターンに差異が見られます。Performerはsoftmaxカーネルを近似しているため、同じスケーリングで比較した場合にMSEが小さくなる傾向があります。

標準Attention Linear Performerの出力比較と近似誤差

実際に生成した図がこちらです。左のパネルでは、PerformerとELU Linear Attentionの出力曲線がどちらも標準Attention(黒)の概形をよくなぞっており、結合順序を変えても出力がきちんと再現できていることがわかります。中央のパネルは位置ごとの近似誤差(対数スケール)で、誤差は $10^{-4}$ 前後と十分小さく抑えられています。右のパネルは標準Attentionの注意重み行列で、この $n \times n$ の行列を明示的に作らずに同等の出力を得ているのがLinear系手法の利点です。

特徴量の次元 $m$ と近似精度の関係

次に、FAVOR+のランダム特徴量の次元 $m$ を変化させたときの近似精度の変化を調べます。

import numpy as np
import matplotlib.pyplot as plt

def favor_plus_feature_map(x, W):
    m = W.shape[0]
    proj = x @ W.T
    norm_sq = np.sum(x ** 2, axis=1, keepdims=True)
    return np.exp(proj - norm_sq / 2) / np.sqrt(m)

def create_orthogonal_random_matrix(m, d):
    n_blocks = (m + d - 1) // d
    blocks = []
    for _ in range(n_blocks):
        H = np.random.randn(d, d)
        Q_mat, _ = np.linalg.qr(H)
        S = np.sqrt(np.random.chisquare(d, size=d))
        blocks.append(S[:, np.newaxis] * Q_mat)
    return np.vstack(blocks)[:m]

def standard_attention_output(Q, K, V):
    d_k = Q.shape[1]
    scores = Q @ K.T / np.sqrt(d_k)
    scores -= np.max(scores, axis=1, keepdims=True)
    e = np.exp(scores)
    return (e / e.sum(axis=1, keepdims=True)) @ V

def performer_output(Q, K, V, m):
    d_k = Q.shape[1]
    Q_s = Q / (d_k ** 0.25)
    K_s = K / (d_k ** 0.25)
    W = create_orthogonal_random_matrix(m, d_k)
    Qp = favor_plus_feature_map(Q_s, W)
    Kp = favor_plus_feature_map(K_s, W)
    KV = Kp.T @ V
    Z = Kp.sum(axis=0)
    num = Qp @ KV
    den = Qp @ Z
    return num / den[:, np.newaxis]

np.random.seed(123)
n, d_k, d_v = 128, 64, 64
Q = np.random.randn(n, d_k) * 0.4
K = np.random.randn(n, d_k) * 0.4
V = np.random.randn(n, d_v) * 0.4

out_std = standard_attention_output(Q, K, V)

# mを変えて近似誤差を測定
m_values = [8, 16, 32, 64, 128, 256, 512]
n_trials = 10  # 各mで複数回試行して平均

mean_mses = []
std_mses = []

for m in m_values:
    trial_mses = []
    for _ in range(n_trials):
        out_perf = performer_output(Q, K, V, m)
        mse = np.mean((out_std - out_perf) ** 2)
        trial_mses.append(mse)
    mean_mses.append(np.mean(trial_mses))
    std_mses.append(np.std(trial_mses))

mean_mses = np.array(mean_mses)
std_mses = np.array(std_mses)

fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# MSE vs m(対数スケール)
axes[0].errorbar(m_values, mean_mses, yerr=std_mses, fmt='o-',
                 capsize=4, linewidth=2, markersize=8, color='#2196F3')
axes[0].set_xscale('log', base=2)
axes[0].set_yscale('log')
axes[0].set_xlabel('Number of random features (m)', fontsize=12)
axes[0].set_ylabel('Mean Squared Error', fontsize=12)
axes[0].set_title('FAVOR+ Approximation Error vs Feature Dim', fontsize=13)
axes[0].grid(True, alpha=0.3)
axes[0].set_xticks(m_values)
axes[0].set_xticklabels([str(m) for m in m_values])

# 1/sqrt(m) の理論的スケーリングと比較
m_theory = np.array(m_values, dtype=float)
scale = mean_mses[3] * np.sqrt(m_values[3])  # m=64を基準に正規化
theory_line = scale / np.sqrt(m_theory)
axes[0].plot(m_values, theory_line, 'r--', linewidth=1.5,
             label=r'$O(1/\sqrt{m})$ scaling', alpha=0.7)
axes[0].legend(fontsize=11)

# m=32 と m=256 の出力比較
out_m32 = performer_output(Q, K, V, m=32)
out_m256 = performer_output(Q, K, V, m=256)

axes[1].plot(out_std[:30, 0], 'k-', label='Standard', linewidth=2)
axes[1].plot(out_m32[:30, 0], 'r--', label='m=32', linewidth=1.5, alpha=0.8)
axes[1].plot(out_m256[:30, 0], 'b--', label='m=256', linewidth=1.5, alpha=0.8)
axes[1].set_xlabel('Position', fontsize=12)
axes[1].set_ylabel('Output (dim 0)', fontsize=12)
axes[1].set_title('Output Comparison: m=32 vs m=256', fontsize=13)
axes[1].legend(fontsize=11)
axes[1].grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig('favor_approximation_error.png', dpi=150, bbox_inches='tight')
plt.show()

for m, mse in zip(m_values, mean_mses):
    print(f"m = {m:>4d}: MSE = {mse:.6f}")

左のグラフは、特徴量の次元 $m$ に対するMSEの変化を対数スケールで示しています。赤い破線は理論的な $O(1/\sqrt{m})$ の収束レートで、実験結果が概ねこの理論曲線に沿って減少していることが確認できます。$m$ を4倍にすると誤差はおよそ半分になるという理論予測と整合しています。右のグラフでは、$m = 32$ と $m = 256$ の出力を標準Attentionと並べて比較しており、$m$ が大きいほど標準Attentionの出力に忠実に近づくことが視覚的にわかります。

エラーバー(各 $m$ で10回試行した結果の標準偏差)を見ると、$m$ が小さいほど試行間のばらつきが大きいことがわかります。これはランダム特徴量のサンプリングに依存する近似手法の宿命であり、直交ランダム特徴量を使うことでこのばらつきを抑えている(独立サンプリングよりは分散が小さい)のがFAVOR+の工夫です。

FAVOR+の近似誤差と特徴量次元mの関係

生成した図の左パネルでは、特徴量の本数 $m$ を8から512まで増やすにつれてMSEが両対数上でほぼ直線的に下がり、赤い破線で示した $O(1/\sqrt{m})$ の理論曲線とよく重なっています。$m$ を4倍にすると誤差がおよそ半分になるという理論予測どおりの挙動です。右パネルでは $m=32$ が標準Attentionからわずかにずれるのに対し、$m=256$ はほぼ完全に重なっており、特徴量を増やすほど近似が忠実になる様子が視覚的に確認できます。

因果的Linear Attentionの再帰実装

RNN的な逐次処理

ここまでの実装は系列全体を行列演算で一括処理するものでしたが、因果的なLinear Attentionは再帰的にも計算できます。自己回帰的なテキスト生成のように、トークンを1つずつ逐次的に処理する場合に有用です。

import numpy as np

def elu_feature_map(x):
    return np.where(x >= 0, x + 1, np.exp(x))

def causal_linear_attention_recurrent(Q, K, V):
    """
    因果的Linear Attentionの再帰実装(RNN的)
    Q: (n, d_k), K: (n, d_k), V: (n, d_v)
    戻り値: (n, d_v)
    """
    n, d_k = Q.shape
    d_v = V.shape[1]

    # 特徴写像の適用
    Q_prime = elu_feature_map(Q)  # (n, d_k)
    K_prime = elu_feature_map(K)  # (n, d_k)

    # 隠れ状態の初期化
    S = np.zeros((d_k, d_v))  # Key-Value要約行列
    z = np.zeros(d_k)          # 正規化ベクトル

    outputs = np.zeros((n, d_v))

    for t in range(n):
        # 状態の更新(加算的)
        S = S + np.outer(K_prime[t], V[t])  # (d_k, d_v)
        z = z + K_prime[t]                   # (d_k,)

        # 出力の計算
        numerator = Q_prime[t] @ S       # (d_v,)
        denominator = Q_prime[t] @ z     # スカラー
        outputs[t] = numerator / denominator

    return outputs

def causal_linear_attention_parallel(Q, K, V):
    """
    因果的Linear Attentionの並列実装(累積和)
    Q: (n, d_k), K: (n, d_k), V: (n, d_v)
    戻り値: (n, d_v)
    """
    Q_prime = elu_feature_map(Q)
    K_prime = elu_feature_map(K)

    n, d_k = Q.shape
    d_v = V.shape[1]

    outputs = np.zeros((n, d_v))

    # 累積和で計算(因果マスクに相当)
    S_cumsum = np.zeros((n, d_k, d_v))
    z_cumsum = np.zeros((n, d_k))

    S_running = np.zeros((d_k, d_v))
    z_running = np.zeros(d_k)

    for t in range(n):
        S_running = S_running + np.outer(K_prime[t], V[t])
        z_running = z_running + K_prime[t]
        S_cumsum[t] = S_running
        z_cumsum[t] = z_running

    for t in range(n):
        num = Q_prime[t] @ S_cumsum[t]
        den = Q_prime[t] @ z_cumsum[t]
        outputs[t] = num / den

    return outputs

# 動作確認
np.random.seed(42)
n, d_k, d_v = 32, 16, 16
Q = np.random.randn(n, d_k) * 0.5
K = np.random.randn(n, d_k) * 0.5
V = np.random.randn(n, d_v) * 0.5

out_rec = causal_linear_attention_recurrent(Q, K, V)
out_par = causal_linear_attention_parallel(Q, K, V)

print(f"Recurrent vs Parallel max diff: {np.max(np.abs(out_rec - out_par)):.2e}")
print("=> 2つの実装は数値的に一致(浮動小数点誤差の範囲)")

再帰実装と並列実装の出力が浮動小数点誤差の範囲で一致することが確認できます。これは、因果的Linear Attentionが再帰的な計算とバッチ的な計算の二つの等価な実装を持つことを意味しています。学習時は並列実装(GPU上で効率的)を使い、推論時は再帰実装(トークンを1つずつ処理)を使うという、状況に応じた使い分けが可能です。

再帰実装のループ内を見ると、各ステップで行っているのは外積の加算 $\bm{S} \leftarrow \bm{S} + \phi(\bm{k}_t)\bm{v}_t^\top$ と行列ベクトル積 $\phi(\bm{q}_t)^\top \bm{S}$ だけです。これはLSTMのゲート計算と構造が似ており、TransformerとRNNの接点を直接的に体感できます。

計算時間の比較実験

系列長を変えた計測

最後に、系列長 $n$ を変化させたときの各手法の計算時間を測定し、理論的な計算量 $O(n^2)$ vs $O(n)$ のスケーリングが実際に観測されるかを確認します。

import numpy as np
import matplotlib.pyplot as plt
import time

def standard_attention_output(Q, K, V):
    d_k = Q.shape[1]
    scores = Q @ K.T / np.sqrt(d_k)
    scores -= np.max(scores, axis=1, keepdims=True)
    e = np.exp(scores)
    return (e / e.sum(axis=1, keepdims=True)) @ V

def elu_feature_map(x):
    return np.where(x >= 0, x + 1, np.exp(x))

def linear_attention_output(Q, K, V):
    Qp = elu_feature_map(Q)
    Kp = elu_feature_map(K)
    KV = Kp.T @ V
    Z = Kp.sum(axis=0)
    num = Qp @ KV
    den = Qp @ Z
    return num / den[:, np.newaxis]

def favor_plus_output(Q, K, V, m):
    d_k = Q.shape[1]
    Q_s = Q / (d_k ** 0.25)
    K_s = K / (d_k ** 0.25)
    n_blocks = (m + d_k - 1) // d_k
    blocks = []
    for _ in range(n_blocks):
        H = np.random.randn(d_k, d_k)
        Qm, _ = np.linalg.qr(H)
        S = np.sqrt(np.random.chisquare(d_k, size=d_k))
        blocks.append(S[:, np.newaxis] * Qm)
    W = np.vstack(blocks)[:m]
    proj_q = Q_s @ W.T
    norm_q = np.sum(Q_s ** 2, axis=1, keepdims=True)
    Qp = np.exp(proj_q - norm_q / 2) / np.sqrt(m)
    proj_k = K_s @ W.T
    norm_k = np.sum(K_s ** 2, axis=1, keepdims=True)
    Kp = np.exp(proj_k - norm_k / 2) / np.sqrt(m)
    KV = Kp.T @ V
    Z = Kp.sum(axis=0)
    num = Qp @ KV
    den = Qp @ Z
    return num / den[:, np.newaxis]

# 計測パラメータ
seq_lengths = [128, 256, 512, 1024, 2048, 4096, 8192]
d_k = 64
d_v = 64
m = 128  # FAVOR+の特徴量次元
n_warmup = 2
n_repeat = 5

times_std = []
times_lin = []
times_perf = []

for n in seq_lengths:
    np.random.seed(0)
    Q = np.random.randn(n, d_k) * 0.3
    K = np.random.randn(n, d_k) * 0.3
    V = np.random.randn(n, d_v) * 0.3

    # ウォームアップ
    for _ in range(n_warmup):
        standard_attention_output(Q, K, V)
        linear_attention_output(Q, K, V)
        favor_plus_output(Q, K, V, m)

    # Standard Attention
    t_list = []
    for _ in range(n_repeat):
        t0 = time.perf_counter()
        standard_attention_output(Q, K, V)
        t_list.append(time.perf_counter() - t0)
    times_std.append(np.median(t_list))

    # Linear Attention (ELU)
    t_list = []
    for _ in range(n_repeat):
        t0 = time.perf_counter()
        linear_attention_output(Q, K, V)
        t_list.append(time.perf_counter() - t0)
    times_lin.append(np.median(t_list))

    # Performer (FAVOR+)
    t_list = []
    for _ in range(n_repeat):
        t0 = time.perf_counter()
        favor_plus_output(Q, K, V, m)
        t_list.append(time.perf_counter() - t0)
    times_perf.append(np.median(t_list))

    print(f"n={n:>5d}: Std={times_std[-1]:.4f}s, "
          f"Linear={times_lin[-1]:.4f}s, Perf={times_perf[-1]:.4f}s")

# 可視化
fig, axes = plt.subplots(1, 2, figsize=(14, 6))

# 絶対時間
axes[0].plot(seq_lengths, times_std, 'o-', label='Standard Attention',
             linewidth=2, markersize=8, color='#E53935')
axes[0].plot(seq_lengths, times_lin, 's-', label='Linear Attention (ELU)',
             linewidth=2, markersize=8, color='#1E88E5')
axes[0].plot(seq_lengths, times_perf, '^-', label='Performer (FAVOR+)',
             linewidth=2, markersize=8, color='#43A047')
axes[0].set_xlabel('Sequence Length (n)', fontsize=12)
axes[0].set_ylabel('Time (seconds)', fontsize=12)
axes[0].set_title('Computation Time vs Sequence Length', fontsize=13)
axes[0].legend(fontsize=11)
axes[0].grid(True, alpha=0.3)

# 対数スケール
axes[1].loglog(seq_lengths, times_std, 'o-', label='Standard ($O(n^2)$)',
               linewidth=2, markersize=8, color='#E53935')
axes[1].loglog(seq_lengths, times_lin, 's-', label='Linear ($O(n)$)',
               linewidth=2, markersize=8, color='#1E88E5')
axes[1].loglog(seq_lengths, times_perf, '^-', label='Performer ($O(n)$)',
               linewidth=2, markersize=8, color='#43A047')

# 理論曲線の参考線
ns = np.array(seq_lengths, dtype=float)
ref_n2 = times_std[2] * (ns / ns[2]) ** 2
ref_n1 = times_lin[2] * (ns / ns[2])
axes[1].loglog(seq_lengths, ref_n2, '--', color='#E53935', alpha=0.4, label=r'$\propto n^2$')
axes[1].loglog(seq_lengths, ref_n1, '--', color='#1E88E5', alpha=0.4, label=r'$\propto n$')

axes[1].set_xlabel('Sequence Length (n)', fontsize=12)
axes[1].set_ylabel('Time (seconds, log scale)', fontsize=12)
axes[1].set_title('Log-Log Scale (Scaling Behavior)', fontsize=13)
axes[1].legend(fontsize=10)
axes[1].grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig('attention_timing_comparison.png', dpi=150, bbox_inches='tight')
plt.show()

左のグラフは系列長に対する計算時間をリニアスケールで示しています。標準Attentionの計算時間が系列長の増加とともに急激に上昇する一方、Linear AttentionとPerformerの計算時間はほぼ線形に増加していることが明確にわかります。$n = 8192$ では、Linear Attentionが標準Attentionより大幅に高速になることが観測できます。

右の対数-対数プロットでは、スケーリングの傾きに注目します。傾き2の破線($O(n^2)$)に沿う標準Attentionと、傾き1の破線($O(n)$)に沿うLinear Attention / Performerの対比が明確に現れます。系列長が短い領域では定数項のオーバーヘッド(特にPerformerではランダム行列の生成コスト)があるため差は小さいですが、系列長が増えるほど線形手法の優位性が拡大していきます。

PerformerはLinear Attention(ELU)よりやや遅いですが、これはランダム行列の生成と射影の追加コストによるものです。GPU上でバッチ処理すればこのオーバーヘッドは相対的に小さくなります。重要なのは、どちらも系列長に対して線形のスケーリングを示している点です。

計算時間と系列長 O(n^2) vs O(n) のスケーリング

実測した計算時間の図がこちらです。左のリニアスケールでは、標準Attention(赤)の計算時間が系列長の増加とともに跳ね上がる一方、LinearとPerformerはほぼ横ばいに近い緩やかな増加にとどまっています。右の両対数プロットでは傾きが本質を語ります。標準Attentionは傾き2の破線($\propto n^2$)に、LinearとPerformerは傾き1の破線($\propto n$)に沿っており、理論どおりのスケーリングが実測でも再現されています。$n=8192$ では標準が約0.64秒なのに対しLinearは約0.02秒と、30倍以上の差がついています。

メモリ使用量の理論的比較

計算時間に加えて、メモリ使用量の違いも確認しておきましょう。

import numpy as np
import matplotlib.pyplot as plt

# メモリ使用量の理論値を計算
seq_lengths = [128, 256, 512, 1024, 2048, 4096, 8192, 16384, 32768, 65536]
d = 64  # モデル次元
m = 128  # ランダム特徴量次元

# Standard Attention: O(n^2) - 注意行列の保持
mem_std = [n ** 2 * 4 / (1024 ** 2) for n in seq_lengths]  # MB (FP32)

# Linear Attention: O(n*d + d*d) - 特徴行列 + 要約行列
mem_lin = [(n * d + d * d) * 4 / (1024 ** 2) for n in seq_lengths]

# Performer: O(n*m + m*d) - 特徴行列 + 要約行列
mem_perf = [(n * m + m * d) * 4 / (1024 ** 2) for n in seq_lengths]

fig, ax = plt.subplots(figsize=(10, 6))

ax.semilogy(seq_lengths, mem_std, 'o-', label='Standard Attention ($O(n^2)$)',
            linewidth=2, markersize=7, color='#E53935')
ax.semilogy(seq_lengths, mem_lin, 's-', label=f'Linear Attention (d={d})',
            linewidth=2, markersize=7, color='#1E88E5')
ax.semilogy(seq_lengths, mem_perf, '^-', label=f'Performer (m={m})',
            linewidth=2, markersize=7, color='#43A047')

# GPU容量の参考線
ax.axhline(y=40 * 1024, color='gray', linestyle=':', alpha=0.5, linewidth=1.5)
ax.text(seq_lengths[0], 40 * 1024 * 1.3, 'A100 40GB', fontsize=10, color='gray')
ax.axhline(y=80 * 1024, color='gray', linestyle=':', alpha=0.5, linewidth=1.5)
ax.text(seq_lengths[0], 80 * 1024 * 1.3, 'A100 80GB', fontsize=10, color='gray')

ax.set_xlabel('Sequence Length (n)', fontsize=12)
ax.set_ylabel('Memory (MB, log scale)', fontsize=12)
ax.set_title('Memory Usage: Attention Matrix Only', fontsize=13)
ax.legend(fontsize=11)
ax.grid(True, alpha=0.3)
ax.set_xscale('log', base=2)
ax.set_xticks(seq_lengths)
ax.set_xticklabels([str(n) for n in seq_lengths], rotation=45)

plt.tight_layout()
plt.savefig('attention_memory_comparison.png', dpi=150, bbox_inches='tight')
plt.show()

print("=== Memory Usage (Attention Matrix Only) ===")
for i, n in enumerate(seq_lengths):
    print(f"n={n:>6d}: Std={mem_std[i]:>10.1f}MB, "
          f"Linear={mem_lin[i]:>8.3f}MB, Perf={mem_perf[i]:>8.3f}MB")

このグラフは、各手法が注意行列の計算に必要とするメモリ量を対数スケールで示しています。標準Attentionのメモリは $n^2$ に比例して急増し、$n = 65536$ では16 GBに達します。これは1ヘッド・1レイヤー分であるため、実際のモデルでは到底GPU1枚に収まりません。

対照的に、Linear AttentionとPerformerのメモリ使用量はほぼ横ばいに見えるほど緩やかに増加しています。$n$ が大きくなっても要約行列 $\bm{\Phi}_K^\top \bm{V} \in \mathbb{R}^{D \times d_v}$ のサイズは系列長に依存しないため(特徴行列 $\bm{\Phi}_Q, \bm{\Phi}_K$ は $O(n)$ で増加しますが $n^2$ に比べれば微小)、超長系列でも実用的なメモリ量で計算が可能です。

注意行列メモリ使用量 O(n^2) vs O(n) の比較

この図は計算量の議論を締めくくる決定的な比較です。標準Attention(赤)のメモリは $n^2$ で急増し、$n=65536$ では16GBに達してGPU 40GB・80GBの破線をやがて超えていきます。一方、LinearとPerformer(青・緑)は系列長を512倍に伸ばしてもメモリがほとんど増えず、対数軸上でほぼ水平のままです。「中間に $n \times n$ 行列を作らない」という設計が、超長系列を扱ううえでいかに本質的かが一目で伝わります。

まとめ

本記事では、Linear AttentionとPerformer(FAVOR+)の理論と実装を解説しました。

  • 標準Attentionの $O(n^2)$ 問題: Self-Attentionのスコア行列 $\bm{Q}\bm{K}^\top$ が $n \times n$ になるため、系列長の二乗に比例する計算量とメモリが必要
  • カーネル解釈: Softmax Attentionをカーネル関数 $\kappa(\bm{q}, \bm{k}) = \exp(\bm{q}^\top \bm{k})$ として再解釈し、特徴写像 $\phi$ で分解する
  • 結合順序の変更: $(\bm{\Phi}_Q \bm{\Phi}_K^\top)\bm{V}$ を $\bm{\Phi}_Q(\bm{\Phi}_K^\top \bm{V})$ に変えるだけで、計算量が $O(n^2)$ から $O(n)$ に劇的に削減される
  • FAVOR+: 正ランダム特徴量と直交ランダム行列を使って、Softmaxカーネルを有限次元で精度よく近似する
  • RNN的解釈: 因果的Linear Attentionは再帰的に計算可能であり、TransformerとRNNを統一的に理解する視点を提供する
  • 実装上のトレードオフ: 特徴量の次元 $m$ を大きくすると近似精度は向上するが計算コストも増加する。$m \approx d$ から $m \approx d \log d$ が実用的な範囲

Linear Attentionは、効率的なシーケンスモデルの理論的な基盤として重要な位置を占めています。近年のState Space Model(S4, Mamba)やRetNetなどは、Linear Attentionの考え方を発展させ、忘却機構やデータ依存のゲーティングを導入したモデルです。Linear Attentionの理論を理解しておくことで、これらの最新モデルの設計原理をより深く理解できるようになります。

次のステップとして、以下の記事も参考にしてください。