每次解碼步驟都會從 GPU HBM 載入整個 KV 快取。在 128K token 時,這會達到數 GB;在 1M token 時則超過 60GB。瓶頸在於記憶體頻寬,而非運算。GPU 必須等待記憶體匯流排傳輸可能與目前查詢無關的資料。
為什麼共享基底稀疏注意力在結構上存在缺陷
Quest、ShadowKV 和 RocketKV 透過估計哪些 KV 頁面重要,並只讀取這些頁面來降低成本。它們透過將目前查詢投影到從完整序列建構的共享低秩基底來對頁面進行評分。根本問題在於:任何固定的基底 W 都存在零空間。如果某個頁面的鍵結構落在該零空間中,無論其對查詢的實際相關性如何,其評分都會接近零。沒有任何訓練程序能解決此問題——這是表示上的不可能。
LOCKS 中的命題 1 將其形式化:每個固定的共享投影都具有「頁面內容盲區方向」。ShadowKV 在小 token 預算時崩潰並非調校問題,而是這種結構性失敗的體現。
LOCKS 建立的觀察
在包含 16 個連續 token 的頁面內,鍵向量會集中在低維子空間中。同時出現的 token 傾向於需要相似的注意力模式,這會產生相似的鍵向量,並在完整的 d_head 維度空間中聚集在小區域內。
在整個序列中,不同頁面佔據不同的子空間——一頁程式碼、一頁散文、一頁引用文獻,每個都具有不同的鍵子空間方向。即使每個個別頁面是局部低秩的,全域鍵矩陣仍然是高秩的。
共享基底必須用單一組方向同時涵蓋所有這些。它無法做到。逐頁基底可以做到,因為每個基底只需涵蓋其自身頁面。
LOCKS 的運作方式
當一個頁面填滿(每 16 個 token)時,LOCKS 會對中心化後的鍵偏差執行一次性 SVD。它儲存:
- μⱼ (int8):頁面質心——16 個 token 的平均鍵向量
- Vⱼ (int4, rank 8):偏差矩陣的前 8 個右奇異向量——此頁面內變異的主方向
- Cⱼ (int8):每個鍵在 Vⱼ 上的投影係數
總計:約為原始頁面 KV 資料的 10%。此操作在解碼關鍵路徑之外執行。
在每個解碼步驟中,LOCKS 僅從這些摘要對每個頁面進行評分——無需讀取 K 或 V 張量:
$$\hat{s}j(q) = \log \sum{i \in P_j} \exp\left[\frac{q^\top \mu_j + (V_j^\top q)^\top c_i}{\sqrt{d}}\right]$$
q^T μⱼ:查詢與頁面平均鍵方向的對齊程度(每個頁面一個純量)。
(Vⱼ^T q)^T cᵢ:查詢與每個個別鍵相對於質心的偏差在頁面局部基底中的對齊程度(每個 token 一個純量)。
透過求和並應用 log-sum-exp,即可得到該頁面的估計總注意力質量。根據此分數選取排名最高的頁面進行完整注意力計算;其餘則跳過。
定理 1 證明未量化秩 r 摘要可達到至少 e^(−4η) 的精確質量選擇覆蓋率,其中 η 界定了 SVD 重建誤差。引理 1 證明在所有確定性值盲選擇器中,質量排序是最壞情況下的最佳策略。
程式碼
import torch
import torch.nn.functional as F
SEQ_LEN = 65536
D_HEAD = 128
PAGE_SIZE = 16
RANK = 8
TOKEN_BUDGET = 2048
K = torch.randn(SEQ_LEN, D_HEAD)
V = torch.randn(SEQ_LEN, D_HEAD)
q = torch.randn(D_HEAD)
# Build page summaries once (off critical path)
summaries = []
for p in range(SEQ_LEN // PAGE_SIZE):
keys = K[p*PAGE_SIZE:(p+1)*PAGE_SIZE]
mu = keys.mean(0)
dev = keys - mu.unsqueeze(0)
_, _, Vh = torch.linalg.svd(dev, full_matrices=False)
basis = Vh[:RANK].T
coeffs = dev @ basis
summaries.append({"mu": mu, "V": basis, "C": coeffs})
# Per-step: score all pages from summaries only
scale = D_HEAD ** -0.5
scores = []
for s in summaries:
global_t = (q @ s["mu"]) * scale
local_t = (s["C"] @ (s["V"].T @ q)) * scale
scores.append(torch.logsumexp(global_t + local_t, dim=0))
page_scores = torch.stack(scores)
# Select top pages and attend
n_top = TOKEN_BUDGET // PAGE_SIZE
top_pages = torch.topk(page_scores, n_top).indices.sort().values
sel_idx = torch.cat([
torch.arange(p*PAGE_SIZE, (p+1)*PAGE_SIZE) for p in top_pages
])
out_locks = F.scaled_dot_product_attention(
q.view(1, 1, D_HEAD),
K[sel_idx].view(1, len(sel_idx), D_HEAD),
V[sel_idx].view(1, len(sel_idx), D_HEAD),
).squeeze()
out_full = F.scaled_dot_product_attention(
q.view(1, 1, D_HEAD),
K.view(1, SEQ_LEN, D_HEAD),
V.view(1, SEQ_LEN, D_HEAD),
).squeeze()
cos_sim = F.cosine_similarity(out_locks.unsqueeze(0), out_full.unsqueeze(0)).item()
print(f"Attended: {len(sel_idx)}/{SEQ_LEN} ({100*len(sel_idx)/SEQ_LEN:.1f}%)")
print(f"Cosine similarity to FullKV: {cos_sim:.4f}")
# Expected: ~0.98+ at 3% token attendance
Enter fullscreen mode Exit fullscreen mode
結果
InfiniteBench (GLM-4-9B-Chat-1M, 100K+ context):
| Method | b=512 | b=2048 |
|---|---|---|
| Quest | 34.7 | 35.7 |
| ShadowKV | 36.0 | 40.3 |
| RocketKV | 38.5 | 42.3 |
| LOCKS | 41.1 | 43.6 |
| FullKV | 43.0 | 43.0 |
在 b=2048 時,LOCKS (43.6) 超越 FullKV (43.0)。對選定的 token 子集進行注意力計算,可排除原本會稀釋注意力分布的無關 token。
LongBench-v1 (Llama-3.1-8B):在 b=64 到 b=2048 的範圍內,與 FullKV 的差距維持在約 1 分以內。其他方法在 b=512 以下即崩潰。
H200 解碼延遲:在 128K 時為 1.59 倍,在 256K 時為 1.80 倍,在 1M 時為 2.0 倍。每步驟讀取的位元組數:在 1M 時減少 9.8 倍。由於在這些上下文長度下,解碼幾乎完全受頻寬限制,頻寬節省幾乎線性地轉換為延遲節省。
注意事項
定理 1 適用於未量化的摘要。已發布的 int4 設定是經過實證驗證,而非每步驟認證。消融實驗顯示 int4 會損失幾個召回率點;int2 則崩潰。投入生產前請進行驗證。
LOCKS 節省的是頻寬,而非 VRAM。完整的 KV 快取仍會常駐。如果限制因素是記憶體容量而非頻寬,則此方法無幫助。
單一作者的預印本;尚未出現獨立複製。在 256K+ 上下文時最有效——在 128K 時 1.59 倍的增益可能不足以證明整合開銷的合理性。
參考文獻
- LOCKS: arXiv:2607.24555
- Quest: arXiv:2406.10774
- ShadowKV: arXiv:2410.21465
- FlashAttention-3: arXiv:2407.08608
0 Comments
Log in to join the conversation.No comments yet. Be the first to share your thoughts.