Skip to content

attention と線形の間

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

attention は系列長の二乗、状態を持つ方式は線形。この2つは別方式に見えて、実は何トークンずつまとめるかというつまみ1つで連続的に繋がっている。出発点は結合則で、softmax を外すと掛ける順を変えられ、K と V を系列長に依らない固定サイズの状態に畳める。チャンクの大きさを両端に振ると、どちらの計算とも値がぴたりと一致する。

この章で作るもの

attentionは全トークン対のスコアを取るので、計算量が系列長の二乗になった。Mamba / SSMは状態を 1 個持って流すので線形になった。本書ではこの 2 つを別の方式として並べてきた。

だが両者は連続的に繋がっている。繋いでいるのは「何トークンずつまとめて処理するか」という数 1 つだけで、それを系列長にすれば attention に、1 にすれば線形になる。間の値も取れる

出発点は行列演算で確かめた結合則だ。行列積は結合的なので (A·B)·CA·(B·C) は同じ値になる。attention がこの性質を使えないのは、間に softmax が挟まっているからだった。外せば使える。

  (q kᵀ) v  ← attention の順          q (kᵀ v)  ← 線形の順

   q      kᵀ                           kᵀ     v
  [L×d] [d×L]                         [d×L] [L×d]
      ↓                                    ↓
   [ L×L ]  ← 系列長の二乗              [ d×d ]  ← 系列長が消える
      × v                                 × q
      ↓                                    ↓
   [ L×d ]                              [ L×d ]   同じ値
掛ける順を変えると、途中に現れるものの大きさが変わる。左は系列長の二乗、右は次元だけで決まる

順に見ていく。

  1. softmax を外すと順を変えられる: 同じ値のまま、途中に現れるものの大きさだけが変わる
  2. チャンクの大きさが両者を繋ぐ: 系列長にすれば attention、1 にすれば線形。値はどこでも同じ
  3. 計算量は固定部と比例部に分かれる: 状態の仕事はつまみに依らず、スコアの仕事だけが比例する

① softmax を外すと順を変えられる

2 つの順序をそのまま書く。行列は整数で持つので、丸めが入らず一致を厳密に確かめられる:

go

// ScoresFirst は (q kᵀ) v の順で計算する。attention の順序。
//
// 先に全対のスコア(L×L)を作るので、そこで系列長の二乗が現れる。
func ScoresFirst(q, k, v Mat) Mat {
	scores := Mul(q, k.T()) // L×L ← ここが二乗
	return Mul(scores, v)
}

// StateFirst は q (kᵀ v) の順で計算する。線形 attention の順序。
//
// 先に kᵀv(d×d)を作るので、系列長に依存しない固定サイズの中間になる。
func StateFirst(q, k, v Mat) Mat {
	state := Mul(k.T(), v) // d×d ← 系列長が出てこない
	return Mul(q, state)
}

// State は K と V から畳んだ状態を返す。大きさは d×d で、系列長に依らない。
func State(k, v Mat) Mat { return Mul(k.T(), v) }

4 トークン・次元 2 の例で両方を計算すると、どちらも [15 17 6 7 15 16 18 20] になった。テストでこの一致を固定している。

変わるのは途中に現れるものだ。

順序途中の形大きさの決まり方
(q kᵀ) vスコア 4×4系列長の二乗
q (kᵀ v)状態 2×2次元だけ。系列長が出てこない

系列長が式から消えるのがこの並べ替えの効果になる。K と V を先に畳んでしまえば、あとは何トークン来ても同じ大きさの箱に足し込むだけで済む。

そのぶん失うものもある。softmax は「どこにどれだけ注目するか」を鋭く決める非線形で、attentionで測ったとおり差を指数で拡大していた。外すと、その鋭さが無くなる。計算の形を得るために、表現の鋭さを手放している

状態の大きさが系列長に依らないことは、そのまま持つ量の差になる。次元 64 で測った。

系列長状態(d×d)KV キャッシュ
1284,09616,384
1,0244,096131,072
8,1924,0961,048,576
65,5364,0968,388,608

系列長を 512 倍にしても状態は 4,096 のまま動かない。テストで、短いうちは状態のほうが大きく、長くなると逆転することも固定した。推論高速化で見た KV キャッシュの伸び方と、ちょうど裏返しになる。

② チャンクの大きさが両者を繋ぐ

ここからが本題になる。系列をチャンクに切り、チャンクの中では全対のスコアを取り、チャンクをまたぐところは状態で運ぶ:

go

// Chunked はチャンクに切って計算する。
//
// チャンクの中は全対のスコア(上三角は落とす)、チャンクをまたぐところは状態で運ぶ。
// C を系列長にすれば CausalScoresFirst と、1 にすれば CausalStateFirst と同じ値になる。
// 途中の C でも結果は変わらず、変わるのは計算の配分だけになる。
func Chunked(q, k, v Mat, c int) Mat {
	if c <= 0 {
		c = 1
	}
	l, d := q.Rows, v.Cols
	out := New(l, d)
	state := New(k.Cols, v.Cols) // d×d。チャンクをまたいで持ち越す

	for start := 0; start < l; start += c {
		end := start + c
		if end > l {
			end = l
		}
		qc, kc, vc := rows(q, start, end), rows(k, start, end), rows(v, start, end)

		// チャンクをまたぐぶんは、持ち越した状態から読む。
		inter := Mul(qc, state)
		// チャンクの中は全対のスコアで混ぜる。自分より後ろは見ない。
		sc := Mul(qc, kc.T())
		tril(&sc)
		intra := Mul(sc, vc)

		for i := start; i < end; i++ {
			for j := 0; j < d; j++ {
				out.Data[i*d+j] = inter.At(i-start, j) + intra.At(i-start, j)
			}
		}
		// 状態を更新して次のチャンクへ持ち越す。
		add(&state, Mul(kc.T(), vc))
	}
	return out
}

func rows(m Mat, from, to int) Mat {
	out := New(to-from, m.Cols)
	copy(out.Data, m.Data[from*m.Cols:to*m.Cols])
	return out
}

func add(dst *Mat, src Mat) {
	for i := range dst.Data {
		dst.Data[i] += src.Data[i]
	}
}

比較の相手として、因果マスクつきの両端も書いておく:

go

// CausalScoresFirst は因果マスクつきで (q kᵀ) v を計算する。attention の側の基準。
//
// 全対のスコアを作ってから上三角を落とすので、L×L の面積をいったん確保する。
func CausalScoresFirst(q, k, v Mat) Mat {
	scores := Mul(q, k.T())
	tril(&scores)
	return Mul(scores, v)
}

// CausalStateFirst は状態を 1 トークンずつ育てながら読む。線形の側の基準。
//
// 位置 i では、そこまでの状態 S_i = Σ_{j≤i} kⱼᵀvⱼ を読む。持つのは d×d だけになる。
func CausalStateFirst(q, k, v Mat) Mat {
	l, d := q.Rows, v.Cols
	out := New(l, d)
	state := New(k.Cols, v.Cols)
	for i := 0; i < l; i++ {
		add(&state, Mul(rows(k, i, i+1).T(), rows(v, i, i+1)))
		o := Mul(rows(q, i, i+1), state)
		copy(out.Data[i*d:(i+1)*d], o.Data)
	}
	return out
}

// tril は上三角(自分より後ろを見る位置)を 0 にする。
func tril(m *Mat) {
	for i := 0; i < m.Rows; i++ {
		for j := i + 1; j < m.Cols; j++ {
			m.Data[i*m.Cols+j] = 0
		}
	}
}

同じ 4 トークンで、チャンクの大きさを 1 から 4 まで振った。

チャンクの大きさ出力
1(線形と同じ)[4 0 1 3 11 8 18 20]
2[4 0 1 3 11 8 18 20]
3[4 0 1 3 11 8 18 20]
4(attention と同じ)[4 0 1 3 11 8 18 20]

どこでも同じ値になる。テストで、C=1 が状態を1つずつ育てる形と、C=L が全対スコアの形と、それぞれ一致することを固定した。割り切れないチャンク(C=3)でも値が変わらないことも確かめている。

つまり attention と線形は、違う答えを出す 2 つの方式ではない。同じ答えに至る道が 2 本あって、その間が連続的に埋まっているだけになる。

書いていて 1 つ間違えた。最初はチャンクの中を非マスクにしていて、チャンクをまたぐ側だけが因果になっていた。これは「どちらでもない」形で、両端のどちらとも一致しない。テストが落ちて気づいた。チャンクの中も外も同じ因果性で揃えて、はじめて連続になる

③ 計算量は固定部と比例部に分かれる

チャンクの大きさを変えても値は同じなら、何が変わるのか。計算量の配分が変わる:

go

// Shape は系列長 L、次元 d、チャンクの大きさ C。
type Shape struct {
	L, D, C int
}

// StateFlops は状態の読み書きにかかる積和の回数。
//
// チャンクごとに「状態へ書く」と「状態から読む」で 2 回、それぞれ C·d² かかる。
// チャンクの本数を掛けると C が消えて 2·L·d² になり、**チャンクの大きさに依らない**。
func (s Shape) StateFlops() int { return 2 * s.L * s.D * s.D }

// ScoreFlops はチャンク内の全対スコアにかかる積和の回数。
//
// チャンクごとに「スコアを作る」と「値を混ぜる」で 2 回、それぞれ C²·d かかる。
// チャンクの本数を掛けると 2·L·C·d になり、**チャンクの大きさに比例する**。
func (s Shape) ScoreFlops() int { return 2 * s.L * s.C * s.D }

// Flops は合計。固定部と、C に比例する部の和になる。
func (s Shape) Flops() int { return s.StateFlops() + s.ScoreFlops() }

// IntraMemory はチャンク内で持つスコア行列の要素数。C×C。
func (s Shape) IntraMemory() int { return s.C * s.C }

// IsFullAttention はチャンクが系列全体と同じ、つまり通常の attention か。
func (s Shape) IsFullAttention() bool { return s.C >= s.L }

// IsLinear はチャンクが 1、つまり線形 attention か。
func (s Shape) IsLinear() bool { return s.C == 1 }

チャンクごとに、状態へ書く仕事とスコアを作る仕事の 2 つがある。チャンクの本数を掛けると、前者からは C が消え、後者には C が残る。系列長 1024、次元 64 で測った。

チャンクの大きさ状態(固定)スコア(比例)合計中間の面積
18,388,608131,0728,519,6801
328,388,6084,194,30412,582,9121,024
648,388,6088,388,60816,777,2164,096
1288,388,60816,777,21625,165,82416,384
2568,388,60833,554,43241,943,04065,536
1,0248,388,608134,217,728142,606,3361,048,576

状態の側は 8,388,608 のまま一度も動かない。動くのはスコアの側だけで、チャンクの大きさに正比例する。テストで、固定部が変わらないこと、比例部が 1024 倍のとき正確に 1024 倍になることを固定した。両端の合計は 16 倍以上開く。

系列長を伸ばしたときの伸び方も分かれる。

系列長C=系列長C=64
25610,485,7604,194,304
51237,748,7368,388,608
1,024142,606,33616,777,216
2,048553,648,12833,554,432

系列長を 2 倍にすると、C=系列長 は約 4 倍(二乗の項が支配する)、C を固定すれば 2 倍のまま(線形)。テストでこの 2 つの伸び方を固定した。

ここから、C は計算量とハードウェアの折り合いをつけるつまみだということになる。FLOPs だけを見れば C=1 が最小だが、行列積の単位が大きい演算器では、小さすぎる行列は性能が出ない。実物が 64 や 128 を使うのは、演算器が効率よく回る粒度がそこにあるからで、計算量の最小と実時間の最小が一致しないという 行列演算で測ったのと同じ話になる。

動かす

下のデモは、チャンクの大きさを動かして値が変わらないことと、計算量の配分が固定部と比例部に分かれる様子を並べて見る。両端に振ると attention と線形のどちらかに一致する。

デモチャンクの大きさで attention と線形がつながるチャンク 1(線形と同じ)
値は変わらない計算量は変わる
チャンク 11234線形と同じ

4 トークン / 次元 2 ・ 整数なので丸めが入らない

40131181820
チャンクを 1 にしても、出力は 4 のときと 1 つも違わない。 状態を1つずつ育てる形 だが、同じ答えに着く

softmax を外すと掛ける順を変えられる。先にスコアを作れば系列長の二乗の面積が要り、先に K と V を畳めば 次元だけで決まる固定サイズの状態になる。チャンクの中は前者、チャンクをまたぐところは後者で運ぶと、 その間を連続的に取れる。

設計の観点

  • 非線形を外すと構造が緩む: softmax があるから結合則が使えなかった。外した瞬間に順序の自由が生まれ、そこから固定サイズの状態が出てくる。何を諦めると何が手に入るかが、この 1 手にまとまっている
  • 同じ答えへの道が複数ある: 値が同じで計算量だけが違うなら、選ぶ基準は正しさではなく資源になる
  • 両端を持つと間が見える: 二乗と線形を別方式として並べている限り、間があることに気づけない。1 つのつまみで書き直して初めて連続だと分かる
  • 固定部は削れない: 状態の仕事はチャンクの大きさに依らない。削れるのは比例部だけなので、下限がある
  • 計算量の最小と速さの最小は別: FLOPs は C=1 が最小だが、実時間は演算器の粒度で決まる
  • 因果性は全体で揃える: 中と外で扱いが違うと、どちらの端とも一致しなくなる。実装で実際に踏んだ

対照と実例

方式チャンクの大きさ途中に持つもの計算量
attention系列長L×L のスコア系列長の二乗
チャンク並列64〜128 が定番C×C のスコア + d×d の状態中間
線形 attention1d×d の状態のみ線形
SSM1(漸化式として書く)固定サイズの状態線形

実例:

  • 線形 attention(2020): softmax を外して特徴写像に置き換え、結合順を変えられるようにした。Mamba / SSMで見た「状態を持って流す」形と、ここで繋がる
  • チャンク並列の実装: 64 や 128 が使われるのは、行列積の演算器が効率よく回る粒度がそこにあるため。FLOPs の最小点とは別のところに実時間の最小点がある
  • ハイブリッド: 層ごとに線形と全対を混ぜる構成が実用に入っている。層という単位で両端を配分する形で、この章のつまみを深さ方向へ広げたものになる

裏どり:

  • 結合則が使えるのは softmax を外したから: 素の attention は softmax(qkᵀ)v で、softmax が行ごとの正規化なので (qkᵀ) を先に確定させないと計算できない。外すことは近似であって、無料の書き換えではない
  • 状態が固定なのは、足し込んでいるから: Σ kᵢᵀvᵢ は何項足しても d×d のまま。だが同じ箱に足し続けるので、容量を超えると過去の情報が互いに干渉する。ここを直す研究(DeltaNetGated DeltaNet)が続いていて、後者は Qwen3.5 などの実モデルに入っている
  • 忘却と名指しの更新は別の能力: 足すだけの状態は上書きができない。特定の 1 つを書き換える仕組みと、まとめて古いものを薄める仕組みは別々に要る、という整理が Gated DeltaNet の出発点になっている
  • FlashAttention は別の軸: あれは同じ二乗の計算を、メモリの読み書きを減らして速くする最適化で、計算量そのものは変わらない。この章のつまみとは直交する
  • この章は正規化を持たない: 実物の線形 attention は分母(スコアの和)で割る。省いたのは結合則の一致を厳密に見せるためで、割り算を入れても順序の話は変わらない

簡略化したこと

  • softmax なし: 外した後の形だけを扱う。特徴写像(ELU+1 など)で近似する話には踏み込まない
  • 正規化なし: 分母で割る処理は省いた
  • 整数行列: 丸めを入れないための選択。実物は浮動小数
  • 1 ヘッド・学習なし: 重みは与えられる前提
  • 状態の干渉は扱わない: 足し込みで過去が混ざる問題と、その対策(デルタ則、ゲート)は裏どりで触れるに留めた
  • 実時間を測っていない: FLOPs は数えるが、演算器の粒度による実速度は測らない

参考資料