Prhub

#50613 [Attention][MLA] Per-request scheduling for MLA chunked context

原始 PR 作者 MatthewBonanni 合并时间 2026-08-07 00:45 文件变更 15 提交数 14 评论 10 代码增减 +940 / -552

执行摘要

MLA prefill 改为按请求调度 chunk,异构 batch 延迟大幅下降

实现 #50497。旧实现中,MLA 预填充把 workspace 平均分给每个带 context 的请求,每个迭代处理所有请求的同一 context 窗口;当 batch 中请求上下文长度差异巨大时,大部分 workspace 被浪费,且每个迭代都要对整个 prefill batch 做 attention 与 merge,chunk 数多、kernel 启动开销大。PR body 明确指出 “MLA prefill chunks are fit into the available workspace rather than forced to be the same size. This can reduce the overall number of chunks and improve prefill latency.”

值得精读。重点关注 plan_mla_context_chunks 的按请求打包与对齐拆分二分搜索(aligned_split_len)、init_mla_context_partial/accumulate_mla_context_chunk 的 partial 生命周期管理,以及如何通过数据结构变更消除额外的 mask kernel。对 MLA、attention 调度和 kernel 合并设计感兴趣的工程师会从中获得有价值的模式。

讨论亮点

Review 中 LucasWilkinson 提出了若干优化建议:将 request_start/request_end 改为 slice 对象(已采纳),质疑 covers_all_context_tokens 早退路径在存在无 context 请求时 attn_output 的形状与初始化是否正确,以及部分测试断言(如 _preserves_request_order)是否过度具体。MatthewBonanni 回应“Yeah I also guarded against having no-context requests”,确认早退路径已做防护。最终 LucasWilkinson 批准:“LGTM thank you!!!”

实现拆解

  1. 重构 metadata 数据结构(vllm/model_executor/layers/attention/mla_attention.py):用 ContextChunk 描述一个 workspace 大小内的连续请求切片,ChunkedContextMetadata 改为包含扁平 chunks 列表、context_lensempty_token_slices,替代原来的按 batch 列的 cu_seq_lens/starts/max_seq_lens 等列表矩阵。
  2. 新增打包调度算法plan_mla_context_chunks 按请求顺序把 context 填进 workspace,超预算时用 aligned_split_len 在块边界对齐地拆分尾部请求;build_mla_chunked_context_metadata 据此生成每个 chunk 完整的 tensor 字段(cu_seq_lensseq_lenstoken_to_seq 等)。这保证“每个 context 行恰好被 gather 一次、chunk 不超 workspace、chunk 内不含空 context 请求”。
  3. 调整 prefill 计算主循环_compute_prefill_context 从按 chunk 索引遍历改为按 chunk 对象遍历,新增 init_mla_context_partialaccumulate_mla_context_chunk 负责 partial 的初始化与增量合并,从而删除了 mask_empty_context 这个额外的 Triton kernel 覆盖步骤。
  4. 统一各 prefill 后端接口:FlashAttention(flash_attn.py)、AITER(aiter_flash_attn.py)、TRTLLM(trtllm_ragged.py)、FlashInfer(flashinfer.py)、tokenspeed(tokenspeed_mla.py)的 run_prefill_context_chunk 签名从 chunk_idx 改为接收 ContextChunk,直接用 chunk 内嵌的 cu_seq_lensmax_seq_len 等,消除了对全局 chunked_context 列表的索引和断言。
  5. 稀疏 MLA 与测试配套:sparse_mla_attention.py 同步遍历 chunks,并将 workspace 下限从 max_num_seqs * block_size 降为 block_size;新增 tests/v1/attention/test_mla_context_chunks.py 覆盖打包不变量、continuation 限定、尾部拆分、DCP 本地 chunk 等;test_mla_backends.py 新增 chunked_context_prefill batch 规格,对所有后端与 SDPA 参考输出对比。
文件 模块 状态 重要度
vllm/model_executor/layers/attention/mla_attention.py 注意力层 modified 9.05
tests/v1/attention/test_mla_context_chunks.py 调度测试 added 8.05
vllm/model_executor/layers/attention/sparse_mla_attention.py 稀疏注意力 modified 7.19
tests/v1/attention/test_mla_backends.py 后端测试 modified 6.92
vllm/v1/attention/ops/triton_merge_attn_states.py 合并算子 modified 6.22

关键符号

plan_mla_context_chunks build_mla_chunked_context_metadata init_mla_context_partial accumulate_mla_context_chunk aligned_split_len run_prefill_context_chunk _compute_context_mha mask_empty_context

关键源码片段

vllm/model_executor/layers/attention/mla_attention.py data-contract

核心实现文件,重写 MLA 分块上下文调度:新增 ContextChunk 数据结构、plan_mla_context_chunks 打包算法、init_mla_context_partial/accumulate_mla_context_chunk,并连带调整 DCP 与稀疏路径。

# _ContextChunkPlan 描述一个 chunk 在请求空间上的布局,
# starts / seq_lens 按请求顺序记录每个请求在 chunk 内的 context 偏移与行数。
@dataclass
class _ContextChunkPlan:
    request_start: int
    request_end: int
    starts: list[int]
    seq_lens: list[int]
    is_continuation: bool
​
​
def plan_mla_context_chunks(
    context_lens: list[int],
    row_budget: int,
    max_context_chunk: int,
    split_alignment: int,
    padded_rows: Callable[[int], int],
) -> list[_ContextChunkPlan]:
    """把各请求的 context 打包到 workspace 大小的 chunk 中。"""
​
    # aligned_split_len 用二分查找在可用行数内选择最大的对齐拆分长度。
    def aligned_split_len(start: int, remaining: int, available: int) -> int:
        max_split = min(max_context_chunk, round_down(remaining - 1, split_alignment))
        low, high = 0, max_split // split_alignment
        while low < high:
            mid = (low + high + 1) // 2
            length = mid * split_alignment
            rows = padded_rows(start + length) - padded_rows(start)
            if rows <= available:
                low = mid
            else:
                high = mid - 1
        return low * split_alignment
​
    plans: list[_ContextChunkPlan] = []
    num_requests = len(context_lens)
    request = 0
    start = 0
    while request < num_requests:
        context_len = context_lens[request]
        if context_len == 0:
            # 无 context 的请求直接跳过,chunk 只覆盖带 context 的请求。
            assert start == 0
            request += 1
            continue
​
        request_start = request
        starts: list[int] = []
        seq_lens: list[int] = []
        rows = 0
        while request < num_requests and context_lens[request] > 0:
            remaining = context_lens[request] - start
            request_rows = padded_rows(start + remaining) - padded_rows(start)
            if rows + request_rows > row_budget:
                # 当前请求放不下时,尝试在块边界对齐拆分尾部。
                split_len = aligned_split_len(start, remaining, row_budget - rows)
                if split_len > 0:
                    starts.append(start)
                    seq_lens.append(split_len)
                    rows += padded_rows(start + split_len) - padded_rows(start)
                    start += split_len
                    break
            rows += request_rows
            starts.append(start)
            seq_lens.append(remaining)
            request += 1
            start = 0
        assert seq_lens
        plans.append(
            _ContextChunkPlan(
                request_start=request_start,
                request_end=request_start + len(seq_lens),
                starts=starts,
                seq_lens=seq_lens,
                is_continuation=starts[0] > 0,
            )
        )
    return plans
tests/v1/attention/test_mla_context_chunks.py test-coverage

新增测试文件,核心验证 chunk 打包不变量(无重叠 / 空洞、不超 workspace、continuation 仅限首请求)以及 DCP 本地 chunk 数据的一致性。

def test_tail_splitting_minimizes_chunks():
    """尾部请求可填满一个 chunk 并在下一个 chunk 头部继续。    若不拆分尾部,这些 context 需要 3 个 chunk;在块边界拆分可降到 2 个,
    即 workspace 行数约束下的最低值。
    """
    metadata = build_chunked_context([768, 512, 512], [3, 5, 7], 1024)
    assert metadata is not None
​
    assert len(metadata.chunks) == 2
    first, second = metadata.chunks
    # 第一个 chunk 覆盖请求 0 的全部 768 行和请求 1 的前 256 行。
    assert first.starts.tolist() == [0, 0]
    assert first.seq_lens.tolist() == [768, 256]
    assert not first.is_continuation
​
    # 第二个 chunk 从请求 1 的偏移 256 继续,并覆盖请求 2 全部 512 行。
    assert second.starts.tolist() == [256, 0]
    assert second.seq_lens.tolist() == [256, 512]
    assert second.is_continuation

评论区精华

init_mla_context_partial 的形状与空 context 请求处理 正确性

LucasWilkinson 询问在 covers_all_context_tokens 提前返回时,如果 batch 中有无 context 的请求,attn_output/attn_softmax_lse 的形状是否正确并已初始化。MatthewBonanni 回应已对无 context 请求做 guard。

结论:确认提前返回路径已处理无 context 请求,形状正确。 · 已解决

covers_all_context_tokens 是否可简化为 len(chunks)==1 设计

LucasWilkinson 建议用 len(chunked_context.chunks) == 1 替代 covers_all_context_tokens,但需确认无 context 请求时形状正确。MatthewBonanni 说明已做额外守卫。

结论:保留辅助方法或等价简化,早退路径安全。 · 已解决

ContextChunk 字段使用 slice 表示范围 设计

LucasWilkinson 建议将 request_start/request_end 和 token_start/token_end 改成 slice 对象,使代码更简洁。

结论:作者采纳,最终使用 request_slice 和 token_slice 字段。 · 已解决

测试是否过度具体化 测试

LucasWilkinson 认为部分断言(如 _preserves_request_order)过于具体,建议移除或放宽。

结论:部分过于具体的断言被精简,保留核心不变量测试。 · 已解决

风险与影响

核心调度路径改动集中在 mla_attention.py,影响所有 MLA 模型(DeepSeek、Kimi 等)的预填充正确性,回归风险高。删除 mask_empty_context 后,空 context 请求的 partial 依赖 init_mla_context_partial 正确初始化;若某些组合未覆盖,merge 可能读到未初始化内存产生 NaN。DCP 路径的本地 chunk 规划(padded_local_*local_starts)复杂度高,跨 rank 一致性问题可能导致 all-gather 结果错位。workspace 下限变化会影响显存占用策略;短请求均匀场景可能因为 per-chunk 字段构建开销略有回退。

对用户:长上下文、异构 batch 的 prefill 延迟显著降低(最坏场景约 4 倍)。对系统:chunk 数量减少降低 kernel 启动与 attention/merge 开销,删除一个 Triton kernel 覆盖 pass 简化了 merge 链路。对团队:MLA prefill 后端接口统一为 ContextChunk,未来新增后端(如 CPU MLA)需适配;同时为 chunked context 调度建立了可验证的不变量测试基线。

核心路径变更 跨后端接口重构 未初始化内存风险 DCP 路径复杂度高 性能收益依赖 batch 异构性

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论