Skip to content

perf: MLA 历史分块 prefill 与稀疏打分优化(默认 16K) - #30

Merged
CURRENTF merged 3 commits into
mainfrom
perf/mla-chunked-prefill
Sep 8, 2026
Merged

perf: MLA 历史分块 prefill 与稀疏打分优化(默认 16K)#30
CURRENTF merged 3 commits into
mainfrom
perf/mla-chunked-prefill

Conversation

@kuma-loong

@kuma-loong kuma-loong commented Sep 8, 2026

Copy link
Copy Markdown
Collaborator

改动解决的问题

MLA 将历史 KV 压缩为 latent 与 RoPE cache,但旧 prefill 会一次 gather 全历史,再展开成逐 head 的 K/V。历史越长,展开张量和投影临时空间越大;单纯提高 workspace 上限只能允许更大的分配,不能限制峰值。

本 PR 将历史的 gather → KV 投影展开 → attention → 合并 整条链按块执行,同时保留稀疏方法所需的观察窗口打分。生产路径收敛为 logits 复用展开 K、probability 扫描 latent,移除四组合实验覆盖开关和旧全量展开入口。

新增配置 mla_prefill_history_chunk_size默认 16384(16K),要求正整数。它控制每次展开的历史 KV token 数,与调度器的 engine_prefill_chunk_sizemax_num_batched_tokens 分别独立;不进行自适应容量搜索或 OOM 后循环重试。中英文配置文档均已同步 16K。

分块 attention 如何执行

  1. 根据当前 batch 的物理 slot 映射和 query packing 构造本步计划,区分每条请求的当前 token 与历史 token。计划按执行 scope 复用,runtime 退役时清理。
  2. 将所有请求的当前 token gather、展开并组成 packed varlen batch,执行一次 causal attention。每条请求的当前 token 仅看到本请求此前的当前 token。
  3. 对每条请求的历史以最多 W 个 token 分块:gather latent/RoPE,执行联合 KV 投影,构造完整 K 和 V,再让该请求的当前 Q 对这个历史块执行 noncausal attention。历史块全部早于当前 Q,所以此处无需块内 causal mask。
  4. 每块返回局部输出与逐 query/head 的 LSE,按 softmax 分母权重合并;块缓冲在消费后释放。完成所有历史块后再转回输出 dtype。

对两个不相交 key 分区,设局部结果为 (O₁, L₁)(O₂, L₂),其中 L = logsumexp(scaled QK)

L = logaddexp(L₁, L₂)
O = exp(L₁ - L) · O₁ + exp(L₂ - L) · O₂

合并采用 FP32 累加,避免每个历史块都进行一次 BF16 舍入。这保持完整历史 attention 的数学语义;不同 GEMM/归约顺序仍可能造成浮点差异,不要求逐 bit 相等。

在支持的 provider 上使用 FA3 contiguous varlen partial attention;没有相应入口时使用绑定的 Triton partial attention 实现。provider 解析不放入 forward 热路径。

内存边界:历史展开量由总历史长度降为单块 W;当前 packed token 仍由调度 token budget 决定。最终 token scores、query/output、LSE 和 metadata 仍占用空间,因此不能将“历史 KV 展开受限”理解成整个 workspace 完全不随输入增长。现有 workspace 预算检查保留,估算按分块 live tensors 计算。

两种稀疏打分方式

记观察 query 集合为 Ω,可参与选择的 causal key 集合为 C(q)。观察窗口只限制使用哪些 Q;probability 仍沿 eligible K 维度归一化,不会沿观察 Q 归一化。

模式 生产计算路径 归约与分母
logits 每块展开 K 后立即复用该 K,在独立评分 kernel 中计算观察 QK 原始 QK 对观察 query 和 head 取最大值;无 softmax 分母
probability 先变换观察 Q,主 attention 完成后分块扫描 latent/RoPE;不再次展开 K/V 沿 eligible causal K 做 softmax,再对观察 query 求平均、head 取最大值

logits:复用主 attention 已展开的 K

score(k) = max over h, q∈Ω of Q[q,h] · K[k,h]

当前 token 与每个历史块都在 K 释放前完成打分。评分 kernel 直接归约出每个 token 的分数,不持久化完整的观察 Q×全部 K 矩阵,也不需要第二遍历史扫描。打分与主 attention 使用同一份已展开 K,可以避免为 logits 额外变换 Q 引入的 BF16 舍入差异。

probability:在 latent 空间重建观察 QK

采用列向量记法,若 K_nope = W_K · c,则:

Q_nopeᵀ · K_nope = (W_Kᵀ · Q_nope)ᵀ · c
z(q,k,h) = scale · [Q_absorbed(q,h)ᵀ · c(k) + Q_rope(q,h)ᵀ · K_rope(k)]
p(q,k,h) = exp(z(q,k,h) - LSE(q,h))
score(k) = max_h mean_{q∈Ω} p(q,k,h)

仅变换观察窗口的 Q;随后扫描共享 latent 与 RoPE cache 计算分数。评分是主 attention 外的独立 kernel,不是向 FA3 内部插入评分逻辑。

分母有两种情况:

  • 全部可见 causal keys 都参与归一化:直接切取主 attention 合并后的最终 LSE,只需一遍 latent 打分扫描。
  • 候选范围排除了 sink/recent 等 key(例如本次 SnapKV):主 attention 的全量 LSE 与候选集合不一致。先扫描候选 latent 累积正确 LSE,再扫描一次计算 probability;两遍都不展开 K/V。跨历史块共享最终分母,不能对每块分别 softmax 后直接相加。

生产没有 latent-logits / 展开-K-probability 的实验选择开关;算法模式仍由 sparse_prefill_score_mode 指定,不自动将 probability 替换成 logits。展开 K 与 absorbed Q 的有限精度计算并非逐 bit 等价,数值测试与 LongBench 分别检查数学误差和最终任务表现。

代码职责与兼容范围

  • MLA prefill operator 负责分块计划、gather/展开/attention/合并和评分执行;底层计算由对应 Triton kernels 实现。
  • MLA attention layer 只连接通用 controller/cache/provider 流程,删除原有全历史 materialization 与重复处理逻辑。
  • PrefillScoreRequest 描述每条请求的观察 query 范围、评分模式与候选边界;cache manager 将现有方法状态映射为物理坐标,operator 消费请求并返回 token_scores。现有方法继续负责评分后的累计、pooling、选择与缓存管理。
  • PrefillComputeView.token_scores 是可选结果交接;没有预计算分数时沿用现有 collector。GQA/MHA 主 attention 与 MLA decode 算法未在本 PR 中改写。
  • FA3 adapter 增加可选 causal/LSE 参数供分块使用;其他调用的默认行为保持原样。删除不再使用的 run_explicit_prefill 及相应 capability 标记。

性能数据:GLM-4.7-Flash / SnapKV

采用新旧实现的配套性能测试结果。H100 80GB、TP=1、BF16、sgl_fa3_sm90,fixed batch=4,输出 512 tokens;观察窗口 32、pool kernel=1、decode keep=2048、结果中的稀疏预算=2176;当前 prefill chunk 与 max batched tokens 均为 8192,GPU utilization=0.95。seed=42,每种长度预热 1 轮、测量 2 轮,每行统计 8 个请求。

已核对新旧所有运行的逐请求 prompt digest、输入/输出长度、iteration 与 request index 一致,全部测量请求成功。16K/32K/64K 为 nominal prompt 长度,实际请求有相同的长度扰动。

新实现的历史块容量扫描

每格为 平均 TTFT(秒) / 端到端输出吞吐(tok/s);W 为历史块容量,不是当前 prefill chunk。

评分 / W Prompt 16K Prompt 32K Prompt 64K
logits / 4K 1.342 / 286.28 3.928 / 180.28 12.113 / 84.44
logits / 8K 1.319 / 286.18 3.596 / 189.56 11.408 / 88.30
logits / 16K(默认) 1.301 / 292.91 3.527 / 197.67 11.066 / 91.09
logits / 32K 1.303 / 283.73 3.532 / 193.26 11.159 / 89.63
probability / 4K 1.376 / 282.92 3.825 / 183.61 12.274 / 83.70
probability / 8K 1.351 / 282.04 3.659 / 188.58 11.542 / 87.85

这组 logits 扫描中 16K 的三个长度 TTFT 与吞吐均优于 4K/8K/32K,支持将默认历史块设为 16K。没有 probability / 16K 实测,不能将 logits 的块容量结论直接推广成 probability 最优值。

同为 W=8K 时,logits 相比 probability 的 TTFT 低约 1%–2%,输出吞吐也接近;新 probability 已不再呈现旧实现中数秒级的额外 TTFT。

与旧实现对比

模式 / Prompt 旧 TTFT → 新 TTFT(秒) TTFT 降低 旧吞吐 → 新吞吐(tok/s)
logits,新 W=16K / 16K 1.331 → 1.301 2.2% 279.73 → 292.91
logits,新 W=16K / 32K 3.624 → 3.527 2.7% 188.07 → 197.67
logits,新 W=16K / 64K 11.441 → 11.066 3.3% 87.81 → 91.09
probability,新 W=8K / 16K 4.052 → 1.351 66.7% 181.39 → 282.04
probability,新 W=8K / 32K 7.402 → 3.659 50.6% 122.06 → 188.58
probability,新 W=8K / 64K 15.267 → 11.542 24.4% 69.81 → 87.85

数据解释:旧结果与新结果使用不同的同型号 H100;因此小幅提升按本批实测报告,不声称严格同卡配对显著性。吞吐使用 output_token_throughput_tps,包含 prefill 对端到端时长的影响,不等同于纯 decode kernel 吞吐或理论 MFU。

probability 的大幅改善来自整条实现路径变化,不能全部归因于 latent 点积更快。旧评分中的长度相关 Triton 编译特化可在新请求长度上触发额外编译;新实现的有界扫描与运行时长度处理也改变了这部分开销。本组端到端记录未分离每个因素的贡献。

LongBench v1:完整 1,500 样本结果

对比新旧实现的 logits 与 probability,共四组完整评估结果。

四组均为 1,500 个 success,无失败;已核对 (dataset, source_idx)、prompt token 数与参考答案一致。repobench-p 为 500 个样本,其余五任务各 200 个。配置记录中的 samples_per_task=20 不代表本次实际计分数量,数量以逐样本文件及 task status 为准。

新质量运行均显式使用 W=8192。其余关键配置:TP=1、batch size=4、max model len=131072、max batched tokens=65536、当前 prefill chunk=8192、window=32、pool=1、decode keep=2048、decode graph 开启、prefix cache 关闭、temperature=0、top_p=1、top_k=1、seed=42、thinking off。

数据集 样本数 旧 logits 新 logits 旧 probability 新 probability
hotpotqa 200 61.63 62.14 61.16 61.25
passage_retrieval_en 200 100.00 100.00 100.00 100.00
repobench-p 500 63.56 63.69 63.43 63.64
qasper 200 39.51 39.61 39.20 38.99
multifieldqa_zh 200 65.12 65.16 65.00 65.19
gov_report 200 31.28 31.30 31.14 31.04
六任务等权平均 1500 60.18 60.32 59.99 60.02
文件中的类别平均 59.20 59.35 58.99 58.98

类别平均未包含 multifieldqa_zh,因此与六任务等权平均不同。新 logits 相比旧 logits 等权平均 +0.13 分;新 probability 相比旧 probability +0.03 分。新 logits 比新 probability 高约 0.30 分,六任务中 4 项领先、1 项持平、1 项略低;这批数据没有显示 probability 的整体质量优势,但未做逐样本显著性检验。

同一 LongBench runtime 配置下,记录的 KV slots 从旧 logits 的 165,635 增至新 logits/probability 的 231,948(+40.0%)。这是这组 W=8K 配置下的可用容量观测,不能直接当作 16K 默认值下的容量,也不代表采样到的整卡显存峰值下降 40%。

基于这些数据,logits 是这组 GLM/SnapKV 工作负载的优先选择;probability 保留其原有算法语义,并获得更低内存的实现。没有新增自动选择打分模式的策略。

验证与复现边界

  • MLA 分块 prefill 与 FA3 的 CUDA 数值/adapter 测试:28 passed。覆盖独立完整 softmax oracle、uneven batch/历史块、slot indirection、FA3/Triton、全量与候选分母、跨多个 query tile、空候选、块容量变化以及评分不重复投影历史 KV。
  • CPU 回归:MLA provider/layer、GLM 模型与稀疏方法、H2O、SnapKV 容量/生命周期,166 passed、3 skipped、4 subtests passed
  • 相关新文件与更新测试通过 Ruff,git diff --check 通过。
  • 以上性能/质量来自提供的完整实验数据,不使用早期 120 样本抽测替代。实验运行对应清理前的开发版本;本 PR 为清理后的生产代码,已有性能与质量结果并非在最终提交上重新执行所得。W=16K 有 logits 性能结果,尚无 W=16K LongBench 或 probability 性能结果。
  • 已核对性能汇总、运行配置、逐请求记录、算子运行信息,以及质量评估的逐样本结果与运行状态。4K logits 运行未显式设置历史块容量,使用当时的 4K 默认值。

@CURRENTF
CURRENTF merged commit 32aefaa into main Sep 8, 2026
1 of 2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants