Skip to content

Mamba / SSMと拡散LM

実装: llm/ssm/ / 実行: go test ./llm/ssm/

attention は全トークンが全トークンを見るので計算量が系列長の二乗になる。状態空間モデル(SSM)は 1 個の状態を系列に沿って更新する線形漸化式で、計算量を系列長に線形化する。ただし固定の更新則では入力を選り分けられない。Mamba の選択的 SSM は状態への取り込み量を入力ごとに変えてこれを解く。線形スキャンとゲート付きスキャンを実装し、attention との計算量の差を確かめる。

この章で作るもの

Transformer全体像 で見たとおり、attention のコストは系列長 n の二乗に比例する。全トークンが全トークンとの内積を取るからで、推論高速化 の KV キャッシュも attention変種の GQA も、この二乗を消すのではなく定数を小さくする工夫だった。二乗そのものを線形に置き換えようとするのが、この章の SSM(state space model)だ。

発想は素朴で、イベントループOS で見た「状態を 1 個持って系列を流す」ループに近い。全ペアを見る代わりに、系列を先頭から 1 つずつ読み、そのたびに 1 個の状態を更新する。1 ステップの計算は定数なので、全体は n に線形になる。問題は、この単純な更新則だと入力を選り分けられないことで、それを解いたのが Mamba だ。

attention (n²)              SSM (n)
 x1 x2 x3 x4                 x1 → x2 → x3 → x4
  ╲│╱╲│╱                     状態 h を持って左から流す
 全ペアの内積                 h ← A·h + B·x を繰り返すだけ
 遠い語も距離1で届く           1ステップ定数・全体で線形
attention と SSM の計算構造。attention は全ペアを見て n²。SSM は 1 個の状態を系列に沿って更新するので n に線形。かわりに全体を一度に見ることはできない

先に押さえることが3つある。

  1. 線形漸化式で状態を運ぶ: h_t = A·h_{t-1} + B·x_t。1 個の状態を更新するだけで系列全体を n に線形で処理する
  2. 固定則の弱点は選択できないこと: どのトークンも同じ規則で状態に入る。関係ある情報だけ残す、ができない
  3. 選択的 SSM が解く: 状態への取り込み量を入力ごとに変えるゲートを足す。これが Mamba の核心

① 線形スキャン: 二次を線形にする

まず基本の SSM を書く。状態 h を持ち、各時刻に線形の更新則で回すだけだ:

go

// SSM は 1 次元の線形状態空間モデル。
//
//	h_t = A·h_{t-1} + B·x_t   (状態の更新)
//	y_t = C·h_t               (状態からの出力)
//
// A が減衰率で、|A|<1 なら過去の影響が指数的に薄れ、A=1 なら状態が保持される。
type SSM struct {
	A, B, C float64
}

// Scan は入力列を先頭から流し、各時刻の出力を返す。
// 1 ステップの計算は定数で、全体は系列長に線形。
func (m *SSM) Scan(x []float64) []float64 {
	y := make([]float64, len(x))
	h := 0.0
	for t, xt := range x {
		h = m.A*h + m.B*xt
		y[t] = m.C * h
	}
	return y
}

// ScanCounted は Scan に加え、状態更新の回数(= 系列長)を返す。
// attention の全ペア計算(AttentionOps)と比べて線形であることを示すための計測点。
func (m *SSM) ScanCounted(x []float64) ([]float64, int) {
	return m.Scan(x), len(x)
}

// AttentionOps は素の attention が長さ n で行う全ペアスコア計算の回数 n²。
// SSM の線形コストと対比するための参照値。
func AttentionOps(n int) int { return n * n }

A が状態の減衰率になる。|A|<1 なら過去の入力の影響は時刻が進むほど指数的に薄れ、A=1 なら状態が保たれて累積和になる。テストでは、先頭に置いたインパルスが減衰率 0.8 で 50 ステップ後にほぼ消えること、A=1 では入力が積み上がることを固定した。

計算量の差は ScanCountedAttentionOps の対比で見える。長さ 1000 の系列で、SSM の状態更新は 1000 回、attention の全ペア計算は 100 万回。n が伸びるほどこの差は開く。100 万トークンのような長系列で SSM が注目される理由がここにある。

RNN との違いも押さえておきたい。見た目は RNN の逐次更新に似ているが、更新則が線形なので、実物は並列プレフィックススキャンでまとめて計算でき、学習も並列に回せる。RNN の逐次性という弱点(Transformer全体像 で attention が置き換えた理由)を、線形性で回避しているのが SSM の要点だ。

② 選択的 SSM: 入力を選り分ける

基本の SSM には致命的な弱点がある。A・B・C が固定なので、どのトークンも同じ規則で状態に入る。attention なら「この語はあの語に強く注目」という入力依存の選択ができるが、固定の漸化式にはそれがない。長い系列で、関係ある情報だけを状態に残して残りを捨てる、という選り分けができないのだ。

Mamba(2023)の選択的 SSM(selective SSM)は、状態への取り込み量を入力ごとに変えるゲートでこれを解く:

go

// Selective は選択的 SSM(Mamba の核)。入力ごとにゲート(取り込み量)が変わる。
// ゲートが 0 に近いトークンは状態にほとんど影響せず素通りし、1 に近いトークンは
// 強く取り込まれる。これで「関係ある情報だけ状態に残す」を系列上で選べる。
type Selective struct {
	decay float64 // 状態の基本減衰率(A に相当)
}

// NewSelective は基本減衰率 decay の選択的 SSM を作る。
func NewSelective(decay float64) *Selective {
	return &Selective{decay: decay}
}

// ScanSelective は入力列 x を、対応するゲート列 gate に従って流す。
// gate[t] は [0,1] にクランプされ、その時刻の入力の取り込み量になる:
//
//	h_t = decay·h_{t-1} + gate_t · x_t
//
// gate=1 は通常の取り込み、gate=0 は入力遮断(状態は減衰のみ)。
func (m *Selective) ScanSelective(x, gate []float64) []float64 {
	y := make([]float64, len(x))
	h := 0.0
	for t := range x {
		g := clamp01(gate[t])
		h = m.decay*h + g*0.1*x[t]
		y[t] = h
	}
	return y
}

func clamp01(v float64) float64 {
	if v < 0 {
		return 0
	}
	if v > 1 {
		return 1
	}
	return v
}

ゲートが 0 に近いトークンは状態にほとんど影響せず素通りし、1 に近いトークンは強く取り込まれる。テストでは、同じ入力列でもゲート列が違えば状態の育ち方が変わること、ゲートを閉じた区間では状態が減衰だけすることを固定した。これで「重要なトークンで状態を更新し、無関係なトークンは無視する」という attention 的な選択が、線形コストのまま手に入る。

実物の Mamba はこれを多次元の状態で行い、ゲートも入力から計算する。素朴な SSM が言語モデルで attention に負けていたのは、まさにこの選択ができなかったからで、選択性を足したことで Transformer に匹敵する品質に届いた。

拡散言語モデル: 生成の順序を変える

SSM が attention の計算構造への挑戦だったのに対し、拡散言語モデルは自己回帰という生成の形への挑戦になる。画像生成の拡散モデル(ノイズから徐々に絵を作る)をテキストに持ち込み、全トークンを同時に、ノイズだらけの状態から何ステップもかけて洗練していく。左から 1 語ずつではないので、生成を並列化でき、後から前を書き直せる。まだ自己回帰モデルの品質には届いていないが、生成の逐次性という制約に正面から挑む数少ない路線として研究が続く。この章では実装せず、方向だけを置いておく。

動かす

下のデモは 2 つの見方を用意した。「計算量」は系列長を伸ばしたとき、SSM の線形コストと attention の二乗コストが開いていく様子をバーで見る。「選択スキャン」は同じ入力列にゲートを与え、開いたトークンだけが状態に効き、閉じた区間で状態が減衰する様子を 1 ステップずつ追える。

デモSSM(線形時間と選択)n = 128
計算量選択スキャン8321285122048
attention(n²)16.4k
SSM(n)128

バーは対数スケール。SSM は 1 個の状態を更新するだけなので、系列長にそのまま線形

系列長 128。attention は全ペアで 16.4k 回、SSM は状態更新 128 回。比は 128 倍で、長いほど開く

SSM は全ペアを見る代わりに 1 個の状態を系列に沿って更新するので、計算が系列長に線形になる。 固定の更新則では入力を選り分けられないが、Mamba の選択的 SSM は取り込み量を入力ごとに変えるゲートで 「重要なトークンだけ状態に残す」を線形コストのまま実現する。

設計の観点

  • 線形 vs 全ペアのトレードオフ: SSM は 1 個の状態に系列を圧縮するので、遠いトークンの情報は状態を経由してしか届かない。attention の「距離 1 で直接見る」表現力とは引き換え。長さで SSM、精密な参照で attention という住み分けになる
  • なぜ純 SSM が主流にならないか: 選択性を足しても、コピーや厳密な検索のような「特定トークンを正確に引く」課題では attention が優る。実物はハイブリッド(層ごとに SSM と attention を混ぜる)に向かっている
  • 推論の効率: SSM は状態が固定サイズなので、生成時に KV キャッシュが要らない。系列が伸びてもメモリが増えないのは長文生成で大きな利点
  • 並列スキャンが鍵: 線形性がなければ RNN と同じ逐次の遅さになる。線形だからこそプレフィックススキャンで学習を並列化できる。「なぜ線形にこだわるか」の答え
  • 拡散 LM という別路線: 上で触れた拡散モデルは、自己回帰(左から 1 語ずつ)自体をやめる方向。SSM が「どう状態を運ぶか」の工夫なのに対し、拡散は「生成の順序」そのものを変える

対照と実例

方式計算量長距離の参照生成時メモリ実例
attentionO(n²)直接(距離1)KVキャッシュが増大GPT、Claude、Llama
SSM(固定)O(n)状態経由のみ固定S4、初期の SSM
選択的 SSMO(n)状態経由 + 選択固定Mamba、Mamba-2
ハイブリッド混在両方中間Jamba(Mamba + attention)
拡散 LM反復ステップ全体を同時研究段階

裏どり:

簡略化したこと

  • 1 次元・スカラー: 実物は状態が多次元、A/B/C は行列。ここはスカラーで漸化式の本質に絞った
  • 並列スキャンなし: 逐次ループで実装。線形性から並列化できることは本文で示すに留めた
  • HiPPO 初期化なし: 長距離記憶に有利な A の初期化は扱わず、減衰の直観だけ
  • ゲートの学習なし: 実物はゲートを入力から学習する。ここは外から与えて挙動を見せた
  • 拡散 LM は実装しない: 生成順序の話として解説に留めた

参考資料