attention と線形の間
実装:
llm/chunked// 実行:go test ./llm/chunked/
attention は系列長の二乗、状態を持つ方式は線形。この2つは別方式に見えて、実は何トークンずつまとめるかというつまみ1つで連続的に繋がっている。出発点は結合則で、softmax を外すと掛ける順を変えられ、K と V を系列長に依らない固定サイズの状態に畳める。チャンクの大きさを両端に振ると、どちらの計算とも値がぴたりと一致する。
この章で作るもの
attentionは全トークン対のスコアを取るので、計算量が系列長の二乗になった。Mamba / SSMは状態を 1 個持って流すので線形になった。本書ではこの 2 つを別の方式として並べてきた。
だが両者は連続的に繋がっている。繋いでいるのは「何トークンずつまとめて処理するか」という数 1 つだけで、それを系列長にすれば attention に、1 にすれば線形になる。間の値も取れる。
出発点は行列演算で確かめた結合則だ。行列積は結合的なので (A·B)·C と A·(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 ] 同じ値順に見ていく。
- softmax を外すと順を変えられる: 同じ値のまま、途中に現れるものの大きさだけが変わる
- チャンクの大きさが両者を繋ぐ: 系列長にすれば attention、1 にすれば線形。値はどこでも同じ
- 計算量は固定部と比例部に分かれる: 状態の仕事はつまみに依らず、スコアの仕事だけが比例する
① softmax を外すと順を変えられる
2 つの順序をそのまま書く。行列は整数で持つので、丸めが入らず一致を厳密に確かめられる:
// 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 キャッシュ |
|---|---|---|
| 128 | 4,096 | 16,384 |
| 1,024 | 4,096 | 131,072 |
| 8,192 | 4,096 | 1,048,576 |
| 65,536 | 4,096 | 8,388,608 |
系列長を 512 倍にしても状態は 4,096 のまま動かない。テストで、短いうちは状態のほうが大きく、長くなると逆転することも固定した。推論高速化で見た KV キャッシュの伸び方と、ちょうど裏返しになる。
② チャンクの大きさが両者を繋ぐ
ここからが本題になる。系列をチャンクに切り、チャンクの中では全対のスコアを取り、チャンクをまたぐところは状態で運ぶ:
// 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]
}
}比較の相手として、因果マスクつきの両端も書いておく:
// 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 つ間違えた。最初はチャンクの中を非マスクにしていて、チャンクをまたぐ側だけが因果になっていた。これは「どちらでもない」形で、両端のどちらとも一致しない。テストが落ちて気づいた。チャンクの中も外も同じ因果性で揃えて、はじめて連続になる。
③ 計算量は固定部と比例部に分かれる
チャンクの大きさを変えても値は同じなら、何が変わるのか。計算量の配分が変わる:
// 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 で測った。
| チャンクの大きさ | 状態(固定) | スコア(比例) | 合計 | 中間の面積 |
|---|---|---|---|---|
| 1 | 8,388,608 | 131,072 | 8,519,680 | 1 |
| 32 | 8,388,608 | 4,194,304 | 12,582,912 | 1,024 |
| 64 | 8,388,608 | 8,388,608 | 16,777,216 | 4,096 |
| 128 | 8,388,608 | 16,777,216 | 25,165,824 | 16,384 |
| 256 | 8,388,608 | 33,554,432 | 41,943,040 | 65,536 |
| 1,024 | 8,388,608 | 134,217,728 | 142,606,336 | 1,048,576 |
状態の側は 8,388,608 のまま一度も動かない。動くのはスコアの側だけで、チャンクの大きさに正比例する。テストで、固定部が変わらないこと、比例部が 1024 倍のとき正確に 1024 倍になることを固定した。両端の合計は 16 倍以上開く。
系列長を伸ばしたときの伸び方も分かれる。
| 系列長 | C=系列長 | C=64 |
|---|---|---|
| 256 | 10,485,760 | 4,194,304 |
| 512 | 37,748,736 | 8,388,608 |
| 1,024 | 142,606,336 | 16,777,216 |
| 2,048 | 553,648,128 | 33,554,432 |
系列長を 2 倍にすると、C=系列長 は約 4 倍(二乗の項が支配する)、C を固定すれば 2 倍のまま(線形)。テストでこの 2 つの伸び方を固定した。
ここから、C は計算量とハードウェアの折り合いをつけるつまみだということになる。FLOPs だけを見れば C=1 が最小だが、行列積の単位が大きい演算器では、小さすぎる行列は性能が出ない。実物が 64 や 128 を使うのは、演算器が効率よく回る粒度がそこにあるからで、計算量の最小と実時間の最小が一致しないという 行列演算で測ったのと同じ話になる。
動かす
下のデモは、チャンクの大きさを動かして値が変わらないことと、計算量の配分が固定部と比例部に分かれる様子を並べて見る。両端に振ると attention と線形のどちらかに一致する。
4 トークン / 次元 2 ・ 整数なので丸めが入らない
softmax を外すと掛ける順を変えられる。先にスコアを作れば系列長の二乗の面積が要り、先に K と V を畳めば 次元だけで決まる固定サイズの状態になる。チャンクの中は前者、チャンクをまたぐところは後者で運ぶと、 その間を連続的に取れる。
設計の観点
- 非線形を外すと構造が緩む: softmax があるから結合則が使えなかった。外した瞬間に順序の自由が生まれ、そこから固定サイズの状態が出てくる。何を諦めると何が手に入るかが、この 1 手にまとまっている
- 同じ答えへの道が複数ある: 値が同じで計算量だけが違うなら、選ぶ基準は正しさではなく資源になる
- 両端を持つと間が見える: 二乗と線形を別方式として並べている限り、間があることに気づけない。1 つのつまみで書き直して初めて連続だと分かる
- 固定部は削れない: 状態の仕事はチャンクの大きさに依らない。削れるのは比例部だけなので、下限がある
- 計算量の最小と速さの最小は別: FLOPs は C=1 が最小だが、実時間は演算器の粒度で決まる
- 因果性は全体で揃える: 中と外で扱いが違うと、どちらの端とも一致しなくなる。実装で実際に踏んだ
対照と実例
| 方式 | チャンクの大きさ | 途中に持つもの | 計算量 |
|---|---|---|---|
| attention | 系列長 | L×L のスコア | 系列長の二乗 |
| チャンク並列 | 64〜128 が定番 | C×C のスコア + d×d の状態 | 中間 |
| 線形 attention | 1 | d×d の状態のみ | 線形 |
| SSM | 1(漸化式として書く) | 固定サイズの状態 | 線形 |
実例:
- 線形 attention(2020): softmax を外して特徴写像に置き換え、結合順を変えられるようにした。Mamba / SSMで見た「状態を持って流す」形と、ここで繋がる
- チャンク並列の実装: 64 や 128 が使われるのは、行列積の演算器が効率よく回る粒度がそこにあるため。FLOPs の最小点とは別のところに実時間の最小点がある
- ハイブリッド: 層ごとに線形と全対を混ぜる構成が実用に入っている。層という単位で両端を配分する形で、この章のつまみを深さ方向へ広げたものになる
裏どり:
- 結合則が使えるのは softmax を外したから: 素の attention は
softmax(qkᵀ)vで、softmax が行ごとの正規化なので(qkᵀ)を先に確定させないと計算できない。外すことは近似であって、無料の書き換えではない - 状態が固定なのは、足し込んでいるから:
Σ kᵢᵀvᵢは何項足しても d×d のまま。だが同じ箱に足し続けるので、容量を超えると過去の情報が互いに干渉する。ここを直す研究(DeltaNet、Gated DeltaNet)が続いていて、後者は Qwen3.5 などの実モデルに入っている - 忘却と名指しの更新は別の能力: 足すだけの状態は上書きができない。特定の 1 つを書き換える仕組みと、まとめて古いものを薄める仕組みは別々に要る、という整理が Gated DeltaNet の出発点になっている
- FlashAttention は別の軸: あれは同じ二乗の計算を、メモリの読み書きを減らして速くする最適化で、計算量そのものは変わらない。この章のつまみとは直交する
- この章は正規化を持たない: 実物の線形 attention は分母(スコアの和)で割る。省いたのは結合則の一致を厳密に見せるためで、割り算を入れても順序の話は変わらない
簡略化したこと
- softmax なし: 外した後の形だけを扱う。特徴写像(ELU+1 など)で近似する話には踏み込まない
- 正規化なし: 分母で割る処理は省いた
- 整数行列: 丸めを入れないための選択。実物は浮動小数
- 1 ヘッド・学習なし: 重みは与えられる前提
- 状態の干渉は扱わない: 足し込みで過去が混ざる問題と、その対策(デルタ則、ゲート)は裏どりで触れるに留めた
- 実時間を測っていない: FLOPs は数えるが、演算器の粒度による実速度は測らない
参考資料
- Katharopoulos et al., Transformers are RNNs(2020) — softmax を外して結合順を変える原典
- Yang et al., Parallelizing Linear Transformers with the Delta Rule(NeurIPS 2024) — チャンク並列と、足し込みの限界への対策
- Yang et al., Gated Delta Networks(ICLR 2025) — 忘却と名指しの更新を組み合わせる
- Dao et al., FlashAttention(2022) — 直交する軸の最適化
- 実装: llm/chunked