「コンテキスト長128kのモデル」と聞くと、長い文章を読める賢いモデル、という印象を持ちます。ところが実際に自分のGPUで動かそうとすると、最初にぶつかるのは賢さの話ではなくメモリの話です。パラメータを全部載せ切ったはずなのに、プロンプトを長くしていくと途中で落ちる。原因はモデルの重みではなく、推論中に溜まっていくKVキャッシュです。
しかも厄介なのは、これが単なる容量の問題ではないことです。自己回帰生成は1トークン出すたびに、そこまでに溜めたKVキャッシュを丸ごと読み直します。トークンを1つ吐く時間が、キャッシュをHBMから読む時間で決まってしまう。つまりKVキャッシュは容量とスループットの両方を同時に握っている、推論の急所です。
この急所に対して、これまでの主流だった答えが MQA / GQA — 「KVヘッドの本数を減らして共有する」方向でした。Multi-head Latent Attention(MLA) はここで発想を変えます。ヘッドは1本も削りません。代わりに、KそのものとVそのものを保存するのをやめて、そこから復元できる細い潜在ベクトルだけを保存するのです。しかも「復元」の計算は、後で見る行列吸収というトリックによって、推論時には実質タダになります。
MLAが効いてくる場面は具体的です。
- 長文脈サービングのコスト: 同じGPUメモリで同時に捌けるリクエスト数(バッチサイズ)が数倍〜数十倍に増える。サービング原価はほぼここで決まります
- 生成スループット: デコードがメモリ帯域律速である以上、読むバイト数が減れば、そのままトークン毎秒が上がる
- 長文脈の実用化: 数十k〜数百kトークンの文脈を、現実的なメモリ量で1台に載せられる
本記事の内容
- なぜKVキャッシュがLLM推論のボトルネックなのか、帯域の観点から数字で押さえる
- MHA → MQA → GQA という「ヘッドを削る」系譜と、その限界
- MLAの低ランク圧縮 — ダウン射影 $\bm{W}^{DKV}$ とアップ射影 $\bm{W}^{UK}, \bm{W}^{UV}$
- 行列吸収(absorption) — Kを実体化せずに済ませるトリックを1行ずつ導出し、数値でも検証する
- RoPEが吸収を壊す理由と、Decoupled RoPE による解決
- クエリ側の低ランク圧縮(学習時メモリの削減)
- MHA / GQA / MQA / MLA のキャッシュ量を記号と具体値で比較
- 低ランクにしても表現力が落ちないのはなぜか、という考察
前提知識
この記事を読む前に、以下の記事を読んでおくと理解が深まります。
- Multi-Head Attentionの理論と実装を完全解説
- KVキャッシュの仕組み — LLM推論を高速化する基本技術
- Multi-Query / Grouped-Query Attention(MQA/GQA)— KVキャッシュを削ってLLM推論を軽くする
- RoPE(回転位置埋め込み)とは?相対位置を回転で表す仕組みを図解でわかりやすく解説
なお本記事では、ベクトルは列ベクトル、行列は左から掛ける規約($\bm{y} = \bm{W}\bm{x}$)で書きます。原論文もこの規約です。
KVキャッシュはなぜ推論の急所なのか
まず「どのくらい重いのか」を、記号ではなく実際の数字で掴んでおきましょう。
1トークン・1層あたりのKVキャッシュ要素数は、素のMulti-Head Attention(MHA)なら
$$ \begin{equation} \text{(1層1トークン)} = 2 \, n_h \, d_h \end{equation} $$
です。$n_h$ がヘッド数、$d_h$ が1ヘッドの次元、先頭の $2$ はKeyとValueの2種類ぶんです。ここに層数 $l$ を掛ければモデル全体になります。
DeepSeek-V2 の公開設定を借りて $l = 60$、$n_h = 128$、$d_h = 128$ とすると、1トークンあたり $2 \times 60 \times 128 \times 128 = 1{,}966{,}080$ 要素。fp16 なら 1トークンで約3.93 MB です。文脈長8192トークンなら 32 GB。これは重み以外に、リクエスト1本ごとに必要な量です。80GBのGPUで、重みを載せた残りに何本入るかを考えると絶望的な数字だとわかります。
さらに深刻なのが帯域です。デコード時、新しいトークンを1つ生成するには、そのステップのQuery1本と、過去すべてのK,V の内積を取ります。過去のK,VはHBMからSRAMへ運ばれます。A100のHBM帯域を約 2 TB/s とすると、32 GB を読むのに
$$ \frac{32\ \text{GB}}{2\ \text{TB/s}} \approx 16\ \text{ms} $$
かかります。1トークン16 ms、つまり約62 トークン/秒が、演算がどれだけ速くても超えられない上限になります(実際には重みの読み出しも別途あるので、さらに下がります)。
ここで演算強度(arithmetic intensity)を見ると事態がはっきりします。キャッシュから読んだK,Vの各要素に対して行う演算は、掛けて足すだけの2 FLOPです。バイトあたり約1 FLOPしかない。一方、現代のGPUはバイトあたり数百FLOPを回せる設計です。デコードは完全にメモリ帯域律速で、GPUの演算器はほとんど遊んでいます。
この状況では、最適化の目標は明確です。FLOPを減らすことではなく、運ぶバイト数を減らすこと。多少計算が増えてもいいから、キャッシュを小さくしたい。これがMQA・GQA・MLAに共通する動機です。
では、これまでどうやってキャッシュを小さくしてきたのでしょうか。系譜を1枚の絵で確認します。
MHA → MQA → GQA — 「ヘッドを削る」という解き方
MHAでは、ヘッド $i$ のQueryはヘッド $i$ のKey・Valueとしか照合しません。1対1対応です。だからKVヘッドの本数はQueryヘッドの本数と同じ $n_h$ 本必要で、キャッシュは $2 n_h d_h$ に膨らみます。
ここで「Queryは人それぞれでいいが、参照する資料は全員で共通でもいいのでは?」と考えたのが MQA(Multi-Query Attention) です。K,Vを1組だけ持って全ヘッドで共有し、キャッシュを $2 d_h$、つまり $1/n_h$ に落とします。極端に効きますが、キー側のヘッドごとの多様性が完全に消えるため、品質の劣化と学習の不安定さが報告されています。
その中間を取ったのが GQA(Grouped-Query Attention) で、$n_g$ 個のグループごとにK,Vを共有します。$n_g = n_h$ ならMHA、$n_g = 1$ ならMQAで、両端を補間するパラメータになっています。実務では $n_g = 8$ あたりがよく使われます。

この図の左3枚は、いずれも「K,Vの箱の本数を減らす」方向の解決だとわかります。1本ずつ独立に持つ(MHA)、4本ずつ束ねる(GQA)、全部1本にする(MQA)。減らせる量は箱の本数で決まるので、限界がはっきりしています。MQAまで行けば $1/n_h$ ですが、そこで打ち止めですし、代償として全ヘッドのキーが文字どおり同一ベクトルになります。
一番右のMLAだけが異質です。箱の本数は減らさず、そもそもK,Vという形で保存するのをやめている。代わりに $\bm{c}^{KV}$ という細い潜在ベクトルを1本だけ持ち、K,Vが必要になったらそこから復元する、という構えです。
ここで当然の疑問が湧きます。復元する計算コストは誰が払うのでしょうか。トークンごとに毎回アップ射影を掛け直すなら、メモリは減っても計算が爆発しそうです。この疑問への答えがMLAの一番おもしろい部分なのですが、まずは圧縮そのものの形を見ておきましょう。
MLAの発想 — 削るのではなく、低ランクに潰す
MLAの出発点は、次の素朴な観察です。
MHAのKVキャッシュは、そもそも情報として冗長ではないか。
トークン $t$ について保存している $2 n_h d_h = 32768$ 個の数値は、どこから来たものでしょうか。すべて、そのトークンの隠れ状態 $\bm{h}_t \in \mathbb{R}^{d}$($d = 5120$)に行列を掛けただけの、決定論的な関数です。$5120$ 次元の入力から作った $32768$ 個の数値なので、独立に動ける自由度は最大でも $5120$ 個しかありません。少なくとも6.4倍は水増しされている、と数え上げだけで言えてしまう。
だったら、その中間にある細い表現のほうを保存すればいい — これがMLAです。数式にすると、ダウン射影で潰してから、必要なときにアップ射影で戻します。
$$ \begin{equation} \bm{c}^{KV}_t = \bm{W}^{DKV} \bm{h}_t, \qquad \bm{c}^{KV}_t \in \mathbb{R}^{d_c} \end{equation} $$
$\bm{W}^{DKV} \in \mathbb{R}^{d_c \times d}$ が「ダウン射影(Down-projection for KV)」です。$d_c$ は潜在次元で、DeepSeek-V2/V3 ではいずれも $d_c = 512$。元の $d = 5120$ から10分の1に絞っています。
そして、KeyとValueはこの潜在から復元します。
$$ \begin{equation} \bm{k}^C_t = \bm{W}^{UK} \bm{c}^{KV}_t, \qquad \bm{v}^C_t = \bm{W}^{UV} \bm{c}^{KV}_t \end{equation} $$
$\bm{W}^{UK}, \bm{W}^{UV} \in \mathbb{R}^{n_h d_h \times d_c}$ が「アップ射影(Up-projection)」です。出力は全ヘッドぶんを連結したもので、ヘッド $i$ の取り分を切り出す行を $\bm{W}^{UK}_i \in \mathbb{R}^{d_h \times d_c}$ と書きます。上付きの $C$ は content(後で出てくるRoPE部と区別するため)を意味します。

図の3段の箱の幅が、そのまま次元数の大きさ(見やすさのため平方根スケール)です。入力 $\bm{h}_t$ は5120次元、そこから作られるK(またはV)は $n_h d_h = 16384$ 次元と入力より太くなっています。MLAが実際にキャッシュに残すのは、その途中にある512次元の細い箱1本だけです。K,V合わせて32768要素を保存していたところが512要素になるので、比にして64倍の圧縮です。
ここで重要なのは、$\bm{W}^{UK}$ が $16384 \times 512$ の縦長の行列だという点です。縦長ということはランクが最大でも512であり、$\bm{k}^C = \bm{W}^{UK}\bm{c}^{KV}$ で作られるKeyは、$16384$ 次元空間の中の512次元の部分空間の中しか動けません。これが「低ランク近似」と呼ばれる所以です(低ランク構造そのものについては特異値分解(SVD)の導出と応用が参考になります)。
さて、先ほど保留にした疑問に戻ります。キャッシュが512要素で済むのはいいとして、Attentionを計算するには結局 $\bm{k}^C$ が要ります。過去 $T$ 個のトークンすべてについて $\bm{W}^{UK}\bm{c}^{KV}_j$ を計算するなら、$T \cdot n_h d_h d_c$ という莫大な演算が毎ステップ発生します。$T = 8192$ なら約690億回の積を1トークンごとに、しかも1層ごとに。これでは本末転倒です。
この計算をまるごと消し去るのが、行列吸収です。
行列吸収 — Kを復元せずに内積する
種明かしをすると、そもそも $\bm{k}^C$ を作る必要がない、というのが答えです。
私たちが本当に欲しいのはKeyそのものではなく、QueryとKeyの内積(Attentionスコア)です。そして内積は結合則が効くので、行列を掛ける順番を自由に組み替えられます。この当たり前の事実が、ここでは決定的に効きます。
ヘッド $i$ のQueryを $\bm{q}_{t,i} = \bm{W}^{Q}_i \bm{h}_t$、Keyを $\bm{k}_{j,i} = \bm{W}^{UK}_i \bm{c}^{KV}_j$ とします。スコアを1行ずつ変形していきましょう。
まず定義通りに代入します。
$$ \begin{equation} \bm{q}_{t,i}^{\top} \bm{k}_{j,i} = \left( \bm{W}^{Q}_i \bm{h}_t \right)^{\top} \left( \bm{W}^{UK}_i \bm{c}^{KV}_j \right) \end{equation} $$
転置の公式 $(\bm{A}\bm{x})^\top = \bm{x}^\top \bm{A}^\top$ を左側に使って、$\bm{h}_t$ を外に出します。
$$ \begin{equation} = \bm{h}_t^{\top} \, {\bm{W}^{Q}_i}^{\top} \, \bm{W}^{UK}_i \, \bm{c}^{KV}_j \end{equation} $$
ここで真ん中の2つの行列に注目します。$\bm{W}^{Q}_i$ も $\bm{W}^{UK}_i$ も学習済みの定数であり、トークン $t$ にも $j$ にも依存しません。ならば、この2つを事前に1枚に潰しておけます。
$$ \begin{equation} \bm{A}_i := {\bm{W}^{Q}_i}^{\top} \bm{W}^{UK}_i \in \mathbb{R}^{d \times d_c} \end{equation} $$
これを使うと、スコアは次のように書けます。
$$ \begin{equation} \bm{q}_{t,i}^{\top} \bm{k}_{j,i} = \bm{h}_t^{\top} \, \bm{A}_i \, \bm{c}^{KV}_j \end{equation} $$
$\bm{k}^C$ という文字がどこにも残っていません。 Keyを一度も作らずにスコアが計算できてしまいました。この式変形が、MLAという手法の核心です。$\bm{W}^{UK}$ がQuery側に「吸収された(absorbed)」という言い方をします。
なお実際のMLAではQuery側も圧縮されているので(後述)、$\bm{h}_t$ は潜在 $\bm{c}^Q_t$ に、$\bm{W}^Q_i$ は $\bm{W}^{UQ}_i$ に置き換わり、$\bm{A}_i \in \mathbb{R}^{d’_c \times d_c}$ になります。議論の形はまったく同じです。
実装上は、この結合をどこで切るかにもう一段の工夫があります。$\bm{A}_i$ をそのまま持つ代わりに、毎ステップ次の1本だけを作ります。
$$ \begin{equation} \tilde{\bm{q}}_{t,i} := {\bm{W}^{UK}_i}^{\top} \bm{q}_{t,i} \in \mathbb{R}^{d_c} \end{equation} $$
すると $\bm{q}_{t,i}^{\top}\bm{k}_{j,i} = \tilde{\bm{q}}_{t,i}^{\top} \bm{c}^{KV}_j$ となります。$\tilde{\bm{q}}$ は現在のトークンにしか依存しないので、1ステップに1回作れば、過去 $T$ 個すべてのキャッシュに対して使い回せます。

上下2つの経路は、数学的にはまったく同じ値を返します。違うのはアップ射影を掛ける場所が「$T$ 個ある側」か「1個しかない側」かだけです。上段の素朴な実装では、キャッシュ内のトークン1個ごとに $n_h d_h d_c = 8{,}388{,}608$ 回の積が要ります。下段の吸収版では、$\tilde{\bm{q}}$ の生成は $T$ に依存せず1回きりで、トークンあたりに残るのは内積の $n_h d_c = 65{,}536$ 回だけ。比にして $d_h = 128$ 倍の削減です。
同じ結合則の付け替えは、Value側にも使えます。
Value側の吸収 — 出力射影に畳み込む
Valueについても、キャッシュから $\bm{v}^C$ を復元したくはありません。ヘッド $i$ の出力を書き下してみます。$a_{tj,i}$ をsoftmax後の重みとして、
$$ \begin{equation} \bm{o}_{t,i} = \sum_j a_{tj,i} \, \bm{v}^C_{j,i} = \sum_j a_{tj,i} \, \bm{W}^{UV}_i \bm{c}^{KV}_j \end{equation} $$
$\bm{W}^{UV}_i$ は $j$ に依存しないので、和の外にくくり出せます。
$$ \begin{equation} \bm{o}_{t,i} = \bm{W}^{UV}_i \left( \sum_j a_{tj,i} \, \bm{c}^{KV}_j \right) \end{equation} $$
つまり潜在空間のまま加重和を取ってから、最後に1回だけアップ射影すればいいわけです。ここでもアップ射影が「$T$ 回」から「1回」に移動しています。
さらに、Attentionの最後には出力射影 $\bm{W}^O$ が待っています。$\bm{W}^O_i \in \mathbb{R}^{d \times d_h}$ をヘッド $i$ の担当ブロックとすると、
$$ \begin{equation} \bm{u}_t = \sum_i \bm{W}^O_i \bm{o}_{t,i} = \sum_i \underbrace{\bm{W}^O_i \bm{W}^{UV}_i}_{=: \, \bm{B}_i \in \mathbb{R}^{d \times d_c}} \left( \sum_j a_{tj,i} \, \bm{c}^{KV}_j \right) \end{equation} $$
$\bm{B}_i$ も学習後は定数なので、モデルロード時に1回作っておけます。$\bm{W}^{UV}$ は $\bm{W}^O$ に吸収されて消えるのです。
こうして、Key側もValue側もアップ射影が推論ループから消えました。キャッシュに残るのは $\bm{c}^{KV}_t$ ただ1本です。ここまでの話は「結合則を付け替えただけ」なので厳密に等価なはずですが、本当にそうか、乱数行列で確かめておきましょう。
数値検証 — 吸収は本当に等価か
小さめの縮小モデルを組んで、素朴な実装と吸収した実装のスコアが一致するかを見ます。
import numpy as np
rng = np.random.default_rng(0)
d, d_c, d_cq, d_h, n_h, T = 256, 64, 96, 32, 4, 6 # 小さめの縮小モデル
H = rng.normal(size=(T, d)) # 各トークンの隠れ状態
W_DKV = rng.normal(size=(d, d_c)) / np.sqrt(d) # KVダウン射影
W_UK = rng.normal(size=(d_c, n_h * d_h)) / np.sqrt(d_c) # Kアップ射影
W_DQ = rng.normal(size=(d, d_cq)) / np.sqrt(d) # Qダウン射影
W_UQ = rng.normal(size=(d_cq, n_h * d_h)) / np.sqrt(d_cq) # Qアップ射影
C_kv = H @ W_DKV # (T, d_c) ← キャッシュに残すのはこれだけ
C_q = H @ W_DQ # (T, d_cq)
# (A) 素朴な実装: 潜在からKを復元し、ヘッドごとに内積する
K = (C_kv @ W_UK).reshape(T, n_h, d_h)
Q = (C_q @ W_UQ).reshape(T, n_h, d_h)
S_naive = np.einsum("tih,jih->tji", Q, K) # (T, T, n_h)
# (B) 吸収した実装: ヘッドごとに W_UQ^T W_UK を1枚の行列へ潰す
A = np.einsum("aih,bih->iab",
W_UQ.reshape(d_cq, n_h, d_h),
W_UK.reshape(d_c, n_h, d_h)) # (n_h, d_cq, d_c)
S_absorb = np.einsum("ta,iab,jb->tji", C_q, A, C_kv)
print("最大絶対誤差:", np.abs(S_naive - S_absorb).max())
print("スコアの代表値:", S_naive[3, 1, 0], S_absorb[3, 1, 0])
print("吸収行列の形:", A.shape, " ランク:", np.linalg.matrix_rank(A[0]))
最大絶対誤差: 9.592326932761353e-14
スコアの代表値: -4.477202187751331 -4.4772021877513355
吸収行列の形: (4, 96, 64) ランク: 32
誤差は $10^{-14}$ 台、これはfloat64の丸め誤差そのもので、両者は数学的に同一です。Kを一切作っていない(B)が、Kを作った(A)と同じスコアを返しています。もう一つ注目してほしいのが最後の行で、吸収行列 $\bm{A}_i$ の形は $96 \times 64$($d_{cq} \times d_c$)なのに、ランクは32、つまり $d_h$ ちょうどです。この事実は後の「表現力」の議論で効いてきます。
続いてValue側です。
W_UV = rng.normal(size=(d_c, n_h * d_h)) / np.sqrt(d_c)
W_O = rng.normal(size=(n_h * d_h, d)) / np.sqrt(n_h * d_h)
P = rng.random((T, T, n_h)); P /= P.sum(axis=1, keepdims=True) # 適当な注意重み
# (A) 素朴: V を復元して加重和 → 連結 → 出力射影
V = (C_kv @ W_UV).reshape(T, n_h, d_h)
out_naive = np.einsum("tji,jih->tih", P, V).reshape(T, n_h * d_h) @ W_O
# (B) 吸収: W_UV を W_O に畳み込み、潜在空間のまま加重和を取る
Bs = [W_UV[:, i*d_h:(i+1)*d_h] @ W_O[i*d_h:(i+1)*d_h, :] for i in range(n_h)]
out_absorb = sum(np.einsum("tj,jb->tb", P[:, :, i], C_kv) @ Bs[i] for i in range(n_h))
print("Value側 最大絶対誤差:", np.abs(out_naive - out_absorb).max())
Value側 最大絶対誤差: 1.1102230246251565e-15
こちらも一致します。$\bm{W}^{UV}$ を $\bm{W}^{O}$ に畳み込んだ $\bm{B}_i$ を使い、潜在ベクトル $\bm{c}^{KV}$ のまま加重和を取っただけで、Valueを復元した場合と同じ出力になりました。Key側・Value側の両方でアップ射影が推論ループから消えることが、これで確認できたことになります。
……と、ここまでは実にきれいな話です。ところが現代のLLMには、この構図を根こそぎ壊す仕掛けが標準装備されています。RoPE です。
RoPEが吸収を壊す
RoPE(回転位置埋め込み)は、位置 $t$ のQueryと位置 $j$ のKeyに、それぞれ位置に応じた回転行列 $\bm{R}_t, \bm{R}_j$ を掛けます。回転行列は直交行列なので $\bm{R}_t^{\top}\bm{R}_j = \bm{R}_{j-t}$ となり、内積が相対位置 $j – t$ にだけ依存するという美しい性質が出てきます。
問題は、その回転がどこに割り込むかです。先ほどの変形をRoPE込みでやり直してみましょう。
$$ \begin{equation} (\bm{R}_t \bm{q}_{t,i})^{\top} (\bm{R}_j \bm{k}_{j,i}) = \bm{h}_t^{\top} \, {\bm{W}^{Q}_i}^{\top} \, \bm{R}_t^{\top} \bm{R}_j \, \bm{W}^{UK}_i \, \bm{c}^{KV}_j \end{equation} $$
$\bm{R}_t^{\top}\bm{R}_j = \bm{R}_{j-t}$ を使うと、真ん中の塊はこうなります。
$$ \begin{equation} \bm{A}_i(j-t) = {\bm{W}^{Q}_i}^{\top} \, \bm{R}_{j-t} \, \bm{W}^{UK}_i \end{equation} $$
$\bm{A}_i$ が相対位置の関数になってしまいました。 定数行列ではないので、事前に1枚へ潰しておくことができません。

上段(RoPEなし)では、$\bm{W}^{Q\top}\bm{W}^{UK}$ という位置に依存しない塊が中央に居座っており、学習が終わった時点で掛け算しておけます。下段(RoPEあり)では、$\bm{R}_{j-t}$ が2つの行列のあいだに挟まってしまい、$j-t$ の値ごとに別の行列になります。文脈長が128kなら理屈のうえで13万枚の行列が必要で、事前計算は不可能です。
$\tilde{\bm{q}}$ を使う実装でも同じことです。$\tilde{\bm{q}}$ が過去トークンごとに違う値になってしまい、「1ステップに1回作って使い回す」という前提が崩れます。結局、キャッシュから $\bm{k}^C_j$ を毎回復元して回転させる羽目になり、潜在ベクトルだけを持つ設計が破綻します。
これを数値でも確認しておきましょう。位置に依存しない1枚の $\bm{A}_i$ でRoPE後のスコアを代用しようとすると、どうなるかを見ます。
def rope(x, pos, base=10000.0):
dd = x.shape[-1]; half = dd // 2
th = pos / (base ** (2 * np.arange(half) / dd))
c, s = np.cos(th), np.sin(th)
x1, x2 = x[..., :half], x[..., half:]
return np.concatenate([x1 * c - x2 * s, x1 * s + x2 * c], axis=-1)
pos = np.arange(T)
Qr = np.stack([rope(Q[t], pos[t]) for t in range(T)]) # content次元に直接RoPEを適用
Kr = np.stack([rope(K[j], pos[j]) for j in range(T)])
S_rope_true = np.einsum("tih,jih->tji", Qr, Kr) # 正しいスコア
S_rope_absorb = np.einsum("ta,iab,jb->tji", C_q, A, C_kv) # 定数1枚で代用
print("RoPE後 真のスコア:", S_rope_true[3, 1, 0],
" 吸収で代用した値:", S_rope_absorb[3, 1, 0])
print("RoPE後 最大絶対誤差:", np.abs(S_rope_true - S_rope_absorb).max())
RoPE後 真のスコア: -2.7349038865409114 吸収で代用した値: -4.4772021877513355
RoPE後 最大絶対誤差: 11.692365022860898
先ほど $10^{-14}$ だった誤差が、いきなり11.7まで跳ね上がりました。値の桁そのものと同じオーダーの誤差なので、「ほぼ合っている」どころか完全に別物です。回転行列は $\bm{q}$ と $\bm{k}$ のあいだに入るため、$\bm{q}$ 側だけを事前に持ち上げても吸収できない、ということが数値からもはっきりします。
低ランク圧縮とRoPEは、このままでは両立しません。どちらかを諦めるのでしょうか。MLAの答えは「分ける」でした。
Decoupled RoPE — 位置を運ぶ次元を別建てにする
発想は単純です。吸収したい部分と、位置情報を運ぶ部分を、別々の次元として持ち、最後に連結する。
ヘッド $i$ のQueryとKeyを、2つのブロックの連結として定義します。
$$ \begin{equation} \bm{q}_{t,i} = \begin{bmatrix} \bm{q}^C_{t,i} \\ \bm{q}^R_{t,i} \end{bmatrix} \in \mathbb{R}^{d_h + d^R_h}, \qquad \bm{k}_{j,i} = \begin{bmatrix} \bm{k}^C_{j,i} \\ \bm{k}^R_{j} \end{bmatrix} \in \mathbb{R}^{d_h + d^R_h} \end{equation} $$
上のブロック(content部)はRoPEを掛けません。これは潜在ベクトルから作られ、先ほどの吸収がそのまま使えます。下のブロック(RoPE部)にだけ回転を掛けます。
$$ \begin{equation} \bm{q}^R_{t,i} = \bm{R}_t \bm{W}^{QR}_i \bm{c}^{Q}_t, \qquad \bm{k}^R_{j} = \bm{R}_j \bm{W}^{KR} \bm{h}_j \end{equation} $$
内積は連結したベクトル同士の内積なので、ブロックごとの和にきれいに分解します。
$$ \begin{equation} \bm{q}_{t,i}^{\top} \bm{k}_{j,i} = \underbrace{(\bm{q}^C_{t,i})^{\top} \bm{k}^C_{j,i}}_{\text{吸収できる}} + \underbrace{(\bm{q}^R_{t,i})^{\top} \bm{k}^R_{j}}_{\text{小さいので素直に持つ}} \end{equation} $$

図の下部にあるとおり、キャッシュは2本立てになります。① 潜在ベクトル $\bm{c}^{KV}_j$(512要素)と、② RoPE済みのキー $\bm{k}^R_j$(64要素)。合計576要素で、$\bm{k}^R$ の追加コストは12.5%に収まっています。
ここで見落としがちな設計判断が2つあります。
第一に、$\bm{k}^R_j$ にはヘッドの添字 $i$ がありません。 全ヘッドで同じ1本を共有します。これは意図的です。もしヘッドごとに持つと $n_h d^R_h = 128 \times 64 = 8192$ 要素になり、潜在ベクトル512要素の16倍という本末転倒なサイズになってしまう。位置情報という「どのヘッドから見てもだいたい同じ意味を持つ情報」だからこそ、共有しても壊れにくいという読みです。なお、キー側が共有でも $\bm{q}^R_{t,i}$ はヘッドごとに別なので、「どの相対位置を重視するか」の重みづけはヘッドごとに自由に学習できます。
第二に、$\bm{k}^R_j$ は潜在 $\bm{c}^{KV}$ からではなく、隠れ状態 $\bm{h}_j$ から直接作ります。 潜在経由にすると、$\bm{c}^{KV}$ が「内容の圧縮」と「位置の運搬」という別種のタスクを兼務することになり、圧縮の効率が落ちます。位置は位置で独立した細い経路を通す、という分離です。
なお $\bm{q}^R$ のほうは潜在 $\bm{c}^{Q}_t$ から作りますが、こちらはキャッシュされないので何次元でもコストに響きません。
softmaxのスケーリングは、連結後の次元 $d_h + d^R_h$ を使って $1/\sqrt{d_h + d^R_h}$ とします(なぜ次元の平方根で割るのかはスケールドドット積の記事を参照してください)。
これで吸収が復活しているか、数値で確かめます。
d_r = 16
W_QR = rng.normal(size=(d_cq, n_h * d_r)) / np.sqrt(d_cq)
W_KR = rng.normal(size=(d, d_r)) / np.sqrt(d) # 全ヘッド共有
QR = (C_q @ W_QR).reshape(T, n_h, d_r)
QR = np.stack([rope(QR[t], pos[t]) for t in range(T)])
KR = H @ W_KR # (T, d_r) ← 2本目のキャッシュ
KR = np.stack([rope(KR[j], pos[j]) for j in range(T)])
# content部は吸収する / RoPE部は素直に内積して足す
S_dec_naive = np.einsum("tih,jih->tji", Q, K) + np.einsum("tir,jr->tji", QR, KR)
S_dec_absorb = np.einsum("ta,iab,jb->tji", C_q, A, C_kv) + np.einsum("tir,jr->tji", QR, KR)
print("Decoupled RoPE 最大絶対誤差:", np.abs(S_dec_naive - S_dec_absorb).max())
Decoupled RoPE 最大絶対誤差: 9.592326932761353e-14
誤差は $10^{-14}$ 台に戻りました。RoPEを入れたにもかかわらず、content部の吸収がそのまま成立しています。位置情報を分離したことで、「圧縮したいものは圧縮でき、位置は位置で運べる」という両立が実現しているわけです。$\bm{k}^R$ のぶんだけキャッシュは増えますが、$d^R_h \ll n_h d_h$ なので誤差のような追加コストで済みます。
ここまではKey・Value側、つまり推論時のキャッシュの話でした。MLAにはもう一つ、目的の異なる低ランク圧縮が入っています。
クエリ側の低ランク圧縮 — 狙いは学習時のメモリ
MLAでは、Queryも同じように潰してから戻します。
$$ \begin{equation} \bm{c}^{Q}_t = \bm{W}^{DQ} \bm{h}_t \in \mathbb{R}^{d’_c}, \qquad \bm{q}^C_t = \bm{W}^{UQ} \bm{c}^{Q}_t \end{equation} $$
DeepSeek-V2/V3 では $d’_c = 1536$ が使われています。KV側の $d_c = 512$ より緩い圧縮です。
ここで、はっきりさせておくべきことがあります。Queryはキャッシュされません。 生成の各ステップで使うのは「今のトークン」のQuery1本だけで、過去のQueryは二度と使わないからです。したがってクエリ側の圧縮は、推論時のKVキャッシュ削減にはまったく寄与しません。
では何のためかというと、学習時です。効くのは主に次の2点です。
- 活性値メモリ: 逆伝播のために順伝播の中間出力を保持する必要があります。素直に $\bm{W}^Q \in \mathbb{R}^{n_h d_h \times d}$ で作るなら、系列長 $\times$ バッチ $\times 16384$ の活性値を層ごとに抱えることになります。$1536$ 次元の中間表現を挟むと、この一部を細くできます
- パラメータ数: $\bm{W}^Q$ 単体なら $16384 \times 5120 \approx 83.9\text{M}$ です。$\bm{W}^{DQ}$ と $\bm{W}^{UQ}$ に分けると $1536 \times 5120 + 16384 \times 1536 \approx 7.9 + 25.2 = 33.1\text{M}$ となり、層あたり約6割減ります。ただしこれは副産物で、狙いはあくまで活性値のほうです
構造としてはLoRAと同じ「太い行列を2枚の細い行列に分解する」形ですが、狙いは違います。LoRAは事前学習済みの重みを凍結したまま差分だけ学習するための分解で、MLAのクエリ圧縮は最初からその形でアーキテクチャを定義しています。
そしてもう一つ、副次的ですが重要な効果があります。$\bm{q}^R_{t,i} = \bm{R}_t \bm{W}^{QR}_i \bm{c}^{Q}_t$ が示すとおり、RoPE部のQueryもこの潜在 $\bm{c}^Q_t$ から作られます。 content部とRoPE部が同じ潜在を共有することで、両者が整合の取れた表現になりやすくなっています。
さて、ここまでで仕組みは出揃いました。では結局どれだけ小さくなるのか、記号と数字の両方で締めておきましょう。
KVキャッシュ容量の比較 — 記号と具体値
1トークン・全 $l$ 層あたりのキャッシュ要素数を、4方式について並べます。
| 方式 | 1トークンあたりのキャッシュ要素数 | 何を保存しているか |
|---|---|---|
| MHA | $2 n_h d_h l$ | 全ヘッドのK,V |
| GQA | $2 n_g d_h l$ | $n_g$ グループぶんのK,V |
| MQA | $2 d_h l$ | 1組だけのK,V |
| MLA | $(d_c + d^R_h) \, l$ | 潜在ベクトルとRoPE用キー |
MLAの行に $2$ の係数が付いていないのが目を引きます。$\bm{c}^{KV}$ 1本からKもVも復元されるので、KeyとValueで同じキャッシュを共有しているからです。KとVを別々に持つ他の3方式との構造的な違いがここに出ています。
DeepSeek-V2 の公開設定 $l = 60,\ n_h = 128,\ d_h = 128,\ d_c = 512,\ d^R_h = 64$ と、比較用に $n_g = 8$ を代入します。
| 方式 | 要素数 / トークン | fp16でのサイズ | MHA比 |
|---|---|---|---|
| MHA | 1,966,080 | 3.932 MB | 100 % |
| GQA ($n_g=8$) | 122,880 | 0.246 MB | 6.25 % |
| MLA | 34,560 | 0.069 MB | 1.76 % |
| MQA | 15,360 | 0.031 MB | 0.78 % |
L, n_h2, d_h2, d_c2, d_r2 = 60, 128, 128, 512, 64
for name, e in [("MHA", 2 * L * n_h2 * d_h2), ("GQA(g=8)", 2 * L * 8 * d_h2),
("MQA", 2 * L * 1 * d_h2), ("MLA", L * (d_c2 + d_r2))]:
print(f"{name:9s} {e:>9,} 要素/token ({e * 2 / 1e6:.3f} MB, fp16)")
MHA 1,966,080 要素/token (3.932 MB, fp16)
GQA(g=8) 122,880 要素/token (0.246 MB, fp16)
MQA 15,360 要素/token (0.031 MB, fp16)
MLA 34,560 要素/token (0.069 MB, fp16)
MLAはMHAの 56.9分の1、実務で広く使われるGQA(8グループ)と比べても 3.56分の1 です。MQAだけはMLAよりさらに小さい(MLAの0.44倍)点は正直に押さえておきましょう。MLAは「最小」ではありません。
MLAが何グループぶんのGQAに相当するかを逆算すると、この位置づけがよく見えます。
$$ \begin{equation} n_g^{\text{equiv}} = \frac{d_c + d^R_h}{2 d_h} = \frac{512 + 64}{2 \times 128} = 2.25 \end{equation} $$
キャッシュ量だけを見れば、MLAは「2.25グループのGQA」と同じです。GQAでグループ数を2まで削ったら、キー側の多様性はかなり失われて品質が落ちます。MLAの主張は、同じキャッシュ量でありながら128ヘッドすべてが別々のキーを持てる、という点にあります。

対数軸であることに注意してください。MHAとMLAの差は2桁近くあります。GQAとMLAの差は見た目では小さく感じますが、3.56倍は「同じGPUに載る同時リクエスト数が3.56倍」という意味なので、サービング原価に直結します。
系列長を伸ばすとどうなるかも見ておきましょう。

全部の線が同じ傾き($T$ に比例)なので、MLAはキャッシュの増え方のオーダーを変えているわけではありません。定数倍を下げているだけです。ただしこの定数倍が効きます。MHAは $T \approx 6100$ でRTX 3090の24 GBを、$T \approx 20000$ でA100の80 GBを、リクエスト1本で使い切ります。MLAなら $T = 131072$ でも約9 GB。同じ24 GBのカードで、MHAが6千トークン1本で音を上げるところを、MLAは13万トークンを2本以上抱えられます。
ここまで、キャッシュが劇的に減ることは確認できました。では品質はどうなのか。512次元まで潰して、本当に情報が失われないのでしょうか。
低ランクでも表現力が落ちないのはなぜか
まず、直感に反するかもしれない事実から始めます。MHAのAttentionスコアも、もともと低ランクの双線形形式です。
ヘッド $i$ のスコアを隠れ状態の言葉で書くと、
$$ \begin{equation} \bm{q}_{t,i}^{\top} \bm{k}_{j,i} = \bm{h}_t^{\top} \left( {\bm{W}^{Q}_i}^{\top} \bm{W}^{K}_i \right) \bm{h}_j \end{equation} $$
中央の行列は $d \times d = 5120 \times 5120$ ですが、$\bm{W}^{Q}_i, \bm{W}^{K}_i \in \mathbb{R}^{d_h \times d}$ の積なので、ランクは高々 $d_h = 128$ です。5120次元の空間で、実質128次元ぶんの情報しか見ていない。MHAは最初から低ランクなのです。
MLAでも同じことが起きます。先ほどの数値検証で $\bm{A}_i$ のランクが $d_h$ ちょうど(コードでは32)になっていたのは、まさにこれでした。$\bm{A}_i = {\bm{W}^{UQ}_i}^{\top}\bm{W}^{UK}_i$ は $d_h$ 次元を経由する積なので、ランクは高々 $d_h$。$d_c \geq d_h$(512 ≥ 128)であるかぎり、ヘッド1本あたりの双線形形式のランクは、MHAとまったく同じです。ここに劣化はありません。
次にキャッシュの情報量を数え直します。MHAが1トークンについて保存する $2 n_h d_h = 32768$ 個の数値は、すべて $\bm{h}_t \in \mathbb{R}^{5120}$ の像です。写像のランクは高々 $5120$。つまりMHAのキャッシュは、情報として最大でも5120次元ぶんしか持っていないのに、32768個の箱に入れて運んでいるわけです。6.4倍の水増しがある。
MLAはこれを $d_c + d^R_h = 576$ まで落とします。$5120$ より小さいので、ここは正真正銘の情報ボトルネックです。それでも性能が保てるということは、「隠れ状態のうちAttentionのK,Vとして本当に使われている情報は、$d = 5120$ よりずっと低次元に収まっている」という経験的事実を意味します。$\bm{W}^{DKV}$ はこの部分空間を学習で見つけているわけです。
では、MQA/GQAとの本質的な違いは何でしょうか。同じ「圧縮」でも、潰す方向が違います。

- MHA(左): ヘッドごとにキーの向きが完全に自由。自由度は最大ですが、キャッシュも最大です
- MQA / GQA(中央): ヘッド間でキーが文字どおり同一のベクトルになります。矢印が1本に潰れている。ヘッドごとに「違うものを見る」という多頭注意の本来の役割が、キー側では失われます
- MLA(右): すべてのヘッドのキーが $d_c$ 次元の共有部分空間に制約されます。しかしその部分空間の中では、$\bm{W}^{UK}_i$ がヘッドごとに違うので、ヘッドごとに別の向きを持てます
制約の掛け方が「ヘッド数を減らす」か「共有部分空間に押し込む」かで違う、というのが図の要点です。MQAは $n_h$ 本の自由度を1本に潰します。MLAは $n_h$ 本の自由度を保ったまま、それらが張る空間の次元を $n_h d_h = 16384$ から $d_c = 512$ に落とします。前者は「ヘッドの個性」という、Transformerの性能に効くことがわかっている構造を直接壊しにいく。後者は、もともと冗長だった次元を削るだけで済む。ここが分かれ目です。
とはいえ「制約なし」ではありません。すべてのヘッドのキーが共通の512次元部分空間に住むので、たとえば「あるヘッドだけが、他のどのヘッドとも直交する方向の特徴を見る」という使い方は、512次元を食い合う形になります。MLAは無料の昼食ではなく、冗長性のうち安全に削れるところを狙って削った設計だと理解するのが正確でしょう。
理論の側はここまでです。最後に、実際に動かすときに必ず出てくる話を整理しておきます。
実装上のトレードオフ
吸収形は「ヘッド次元576のMQA」になる。 吸収した後のAttentionを眺めると、面白い見え方ができます。キャッシュ側は $\bm{c}^{KV}_j$(512次元)と $\bm{k}^R_j$(64次元)を連結した576次元のベクトル1本、Query側はヘッドごとの $\tilde{\bm{q}}_{t,i}$ と $\bm{q}^R_{t,i}$ を連結した576次元。つまりキー1本を全ヘッドで共有する、ヘッド次元576のMQAとまったく同じ形です。MLA専用カーネルがMQA系のカーネル構造を土台に書かれるのは、このためです。
FLOPは増える。 吸収形では、キャッシュのトークン1個あたり $n_h d_c = 65{,}536$ 回の積が要ります。MHAの $n_h d_h = 16{,}384$ の4倍です。MLAは演算を減らしていません。メモリ転送量を57分の1にする代わりに、演算を4倍払っているのです。デコードがメモリ帯域律速である以上これは大幅に得な取引ですが、条件が変われば話も変わります。
プリフィルとデコードで最適な形が違う。 長いプロンプトを一気に処理するプリフィル時は、そもそも計算律速です。この局面では吸収せず、素直に $\bm{k}^C, \bm{v}^C$ を実体化して普通のMHAとして計算するほうが速くなります(アップ射影のコストは全トークンにわたって行列積として償却できるため)。逆にデコードは吸収形が有利。実装によっては、この2つを局面ごとに切り替えます。
標準のFlash Attentionカーネルはそのままでは使えない。 一般的なFlash Attention実装は、ヘッド次元が64や128であることを前提にチューニングされています。576という中途半端に大きな次元、しかもKとVが同一ベクトルという構造は想定外なので、MLA専用のカーネルが別途必要になります。Flash AttentionやPagedAttentionといった既存の最適化と組み合わせる場合、実装の追従が必要な部分です。
正規化の扱い。 実装では潜在ベクトルにRMSNormが挟まります。RMSNormはトークンごとのスカラー除算と学習された対角ゲインの積なので、正規化後の潜在をキャッシュすれば、対角ゲインは $\bm{W}^{UK}$ 側に畳み込めて吸収は保たれます。正規化前を保存すると、トークンごとのスカラーを別途持つ必要が出てきます。
既存モデルからの乗り換えは容易ではない。 GQAで学習済みのモデルにMLAを後付けしようとすると、$\bm{W}^{K}, \bm{W}^{V}$ をSVDなどで低ランク分解して $\bm{W}^{DKV}, \bm{W}^{UK}, \bm{W}^{UV}$ に変換したうえで、RoPE部を切り出し直す作業が要ります。近似が入るため追加学習が前提になります。GQAのように「MHAから少ない追加学習で作り替える(uptraining)」ほど手軽ではありません。
MoEとの組み合わせ。 MLAを採用したDeepSeek-V2/V3は、FFN側にMixture of Expertsを併用しています。MoEは「重みは巨大だが1トークンあたりの計算は少ない」構造で、MLAは「キャッシュを小さくする」構造です。役割が重複しておらず、推論コストの別々の項を叩いている点で、相性のよい組み合わせだと言えます。
まとめ
本記事では、Multi-head Latent Attention(MLA)について解説しました。
- 動機: LLMのデコードはメモリ帯域律速で、1トークンごとにKVキャッシュを丸ごと読み直す。キャッシュ削減はそのままスループット改善になる
- 系譜との違い: MHA → MQA → GQA は「KVヘッドの本数を削って共有する」方向だった。MLAはヘッドを削らず、K,Vを低ランクの潜在ベクトル $\bm{c}^{KV}_t = \bm{W}^{DKV}\bm{h}_t$ に圧縮して保存する
- 行列吸収: $\bm{q}^\top\bm{k} = \bm{h}_t^\top ({\bm{W}^{Q}}^{\top}\bm{W}^{UK}) \bm{c}^{KV}$ と結合を付け替えれば、$\bm{k}^C$ を一度も作らずにスコアが計算できる。$\bm{W}^{UV}$ も同様に $\bm{W}^O$ へ吸収される。数値検証でも誤差 $10^{-14}$ 台で一致
- RoPE非互換: 回転行列が $\bm{q}$ と $\bm{k}$ のあいだに挟まると、吸収行列が相対位置の関数になって事前計算できない。数値でも誤差が11.7まで跳ね上がる
- Decoupled RoPE: 位置を運ぶ $d^R_h = 64$ 次元だけを別建てにし、全ヘッドで共有して素直にキャッシュする。content部の吸収は保たれる
- キャッシュ量: MLAは $(d_c + d^R_h)\,l$。DeepSeek-V2構成でMHAの1.76%、GQA(8)の3.56分の1。GQAでいえば2.25グループ相当のキャッシュ量で、128ヘッド全部が別々のキーを持てる
- 表現力: ヘッド1本あたりの双線形形式のランクは $d_h$ のままで、MHAと変わらない。MQA/GQAが「ヘッドを潰す」のに対し、MLAは「共有部分空間に押し込む」— 潰す方向が違う
- 代償: 演算量はMHAの4倍。専用カーネルが必要で、既存モデルからの変換も容易ではない
MLAの教訓を一般化するなら、「キャッシュとは、後で使う値を保存する場所ではなく、後で復元できる情報を保存する場所である」ということに尽きます。復元コストが結合則の付け替えでゼロにできるなら、保存すべきは結果ではなく種のほうです。この視点は、KVキャッシュの量子化やスパース化といった他の削減手法を考えるときにも効いてきます。
次のステップとして、以下の記事も参考にしてください。