每次解码步骤都会从 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,秩 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 得到该页的估计总注意力质量。按此分数选择前 k 页进行完整注意力;其余跳过。

定理 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+ 上下文):

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× 增益可能不足以抵消集成开销。


参考文献