Prhub

#46202 [CPU] Enable chunked prefill and prefix caching for qwen3.5

原始 PR 作者 tianmu-li 合并时间 2026-06-25 11:49 文件变更 7 提交数 17 评论 6 代码增减 +388 / -32

执行摘要

CPU 线性注意力模型支持 chunked prefill 和前缀缓存

之前CPU平台上对线性注意力模型无条件下禁用chunked prefill和prefix caching(见vllm/platforms/cpu.py),因为担心正确性问题。经过验证底层内核修复(conv忽略has_initial_state)和模型运行器调整,这些功能现在可以安全启用,从而提升CPU推理吞吐并降低首token延迟。PR body指出目的为移除过度保守的guard并修复混合模型中部分写入块的问题。

  • 值得精读 _zero_block_ids 的实现,它参考 GPU KVBlockZeroer 进行适配,是混合注意力模型正确性的关键。
  • batch_memcpy 回退展示了如何在缺乏 Triton 时用 ctypes 兼容核心特性,值得其他后端参考。
  • conv.cpphas_initial_state 修复是 chunked prefill 正确的核心。
  • 建议关注 depthfirst-app[bot] 提出的空指针问题是否在后续修复。
讨论亮点
  • 正确性验证:Reviewer @bigPYJ1151 提出“Did you check the accuracy?”,作者 @tianmu-li 回复已验证,并在调节 max_num_batched_tokens=128 后发现精度问题后修复,最终确认 chunked prefill 与 full prefill 的 logprobs 接近(gsm8k 准确率 0.92)。
  • 空指针风险:自动化检查工具 depthfirst-app[bot] 在 csrc/cpu/sgl-kernels/conv.cpp 指出:新代码仅在 has_conv_states 真时访问 has_initial_state[bs],但 has_initial_state 可能是空指针,建议增加 has_initial_state != nullptr 检查。该建议尚未看到对应修复提交。

实现拆解

  1. 移除vllm/platforms/cpu.py中的禁用Guard
    CpuPlatform.check_and_update_config 中删除了根据 AMX 支持和线性注意力层数量强制设置 enable_prefix_caching=Falseenable_chunked_prefill=False 的代码块。

  2. 修复CPUModelRunner._zero_block_ids
    从空操作改为对 FullAttentionSpec 类型的 KV 缓存块数据清零,避免混合注意力模型中非全注意力块的部分写入导致脏数据影响计算。使用 data_ptr 去重防止重复清零。

  3. 修复C++ conv内核忽略has_initial_state
    csrc/cpu/sgl-kernels/conv.cppcausal_conv1d_fwd_varlen_kernel_impl 中,将硬编码的 nullptrfalse 替换为从参数传递的 conv_stateshas_initial_state[bs],使 varlen 卷积在继续 prefill chunk 时能正确使用前一步的卷积状态。

  4. 添加batch_memcpy_kernel的CPU Fallback
    vllm/utils/cpu_triton_utils.py 中实现 _batch_memcpy_impl,基于 ctypes.memmove 批量拷贝内存,用于 triton-cpu 不可用时替代 mamba_utils.batch_memcpy_kernel。在 cpu_model_runner._postprocess_triton 中注册该回退。

  5. 添加回归测试
    - 在 tests/kernels/mamba/cpu/test_cpu_gdn_ops.py 中新增 test_chunk_gated_delta_rule_cpu_two_call_split 等测试,模拟两次调度步骤间的状态传递,验证内核跨 chunk 的正确性。
    - 在 tests/v1/e2e/test_cpu_linear_attn_chunked_prefix.py 中新增端到端测试,使用 check_logprobs_close 比较 chunked prefill 与 full prefill 的 logprobs,以及前缀缓存命中与冷缓存的输出。
    - CI 配置调整:增大 CPU 语言/池化测试超时,将线性注意力测试移动至 triton-cpu job。

文件 模块 状态 重要度
vllm/platforms/cpu.py 平台层 modified 6.31
vllm/v1/worker/cpu_model_runner.py 推理引擎 modified 6.91
vllm/utils/cpu_triton_utils.py 工具库 modified 6.31
csrc/cpu/sgl-kernels/conv.cpp CPU 内核 modified 5.89
tests/kernels/mamba/cpu/test_cpu_gdn_ops.py GDN 算子 modified 7.2
tests/v1/e2e/test_cpu_linear_attn_chunked_prefix.py 线性注意力 added 7.2

关键符号

check_and_update_config _zero_block_ids _batch_memcpy_impl batch_memcpy_kernel causal_conv1d_fwd_varlen_kernel_impl test_chunk_gated_delta_rule_cpu_two_call_split test_causal_conv1d_torch_two_call_split test_causal_conv1d_fwd_cpu_two_call_split test_batch_memcpy_cpu_fallback test_chunked_prefill_matches_full_prefill test_prefix_cache_hit_matches_cold_cache

关键源码片段

vllm/platforms/cpu.py core-logic

移除强制禁用 chunked prefill 和 prefix caching 的 guard,这是功能启用的核心开关。

# vllm/platforms/cpu.py (check_and_update_config excerpt)
# 下面的 if 块已在 PR#46202 中删除:
# if torch.cpu._is_amx_tile_supported() and (
# model_config is not None
# and model_config.get_num_layers_by_block_type(
# parallel_config, \"linear_attention\"
# ) > 0
# ):
# cache_config.enable_prefix_caching = False
# scheduler_config.enable_chunked_prefill = False
# logger.warning_once(
# \"Disabled unsupported prefix caching and chunked prefill \"
# \"for linear attention on AMX CPU platforms.\"
# )
# 现在这两个配置项可以根据用户设置正常使用。
vllm/v1/worker/cpu_model_runner.py core-logic

修复 _zero_block_ids 对混合注意力模型中非全注意力块的部分写入问题,是正确性关键。

def _zero_block_ids(self, block_ids: list[int]) -> None:
    # Zero full-attention blocks to prevent stale data corruption on partial writes.
    # `FullAttentionSpec` 过滤排除了 encoder-only 层,避免误清零
    seen_ptrs: set[int] = set()
    for group in self.kv_cache_config.kv_cache_groups:
        # 只处理 FullAttentionSpec 的 KV 缓存组
        if not isinstance(group.kv_cache_spec, FullAttentionSpec):
            continue
        for layer_name in group.layer_names:
            ctx = self.compilation_config.static_forward_context.get(layer_name)
            if ctx is None:
                continue
            kv = ctx.kv_cache
            if not isinstance(kv, torch.Tensor):
                continue
            # 避免重复清零同一 Tensor
            if kv.data_ptr() in seen_ptrs:
                continue
            seen_ptrs.add(kv.data_ptr())
            for block_id in block_ids:
                kv[block_id].zero_() # 清零块数据

评论区精华

正确性验证 正确性

Reviewer bigPYJ1151 询问是否检查了 accuracy,作者 tianmu-li 回复已运行 gsm8k 基准测试并修复了 max_num_batched_tokens=128 时的精度问题,最终确认 chunked prefill 与 full prefill 的 logprobs 接近。

结论:解决,准确率 0.92,与 full prefill 一致。 · 已解决

has_initial_state 空指针风险 正确性

depthfirst-app[bot] 指出 `has_initial_state` 可能为空指针但解引用时未检查,建议增加 `has_initial_state != nullptr` 检查。

结论:未在 PR 中修复,可能需要后续 PR 处理。 · unresolved

风险与影响

  • 空指针风险:在 csrc/cpu/sgl-kernels/conv.cppcausal_conv1d_fwd_varlen_kernel_impl 中,新代码 has_initial_states_value = has_conv_states ? has_initial_state[bs] : false 仅检查 has_conv_states 而不检查 has_initial_state 是否为空指针;如果 conv_states 非空但 has_initial_state 为空,则导致解引用空指针。当前调用点可能总是同时传参,但缺乏防御性检查。
  • 部分写入脏数据_zero_block_ids 的修复解决了混合注意力模型在 chunked prefill 时的脏数据问题,但仅对 FullAttentionSpec 类型清零;若未来有其他注意力规格,可能需要更新。
  • 性能回退batch_memcpy 使用 Python 循环加 ctypes.memmove 实现,性能远低于原生 Triton 内核,但在 triton-cpu 不可用时的 fallback 路径触发,影响前缀缓存命中的 Mamba 状态拷贝,频率不高。
  • 测试覆盖有限:端到端测试仅覆盖 Qwen3.5-0.8B 单模型和有限 prompt 长度,未涵盖所有线性注意力变体或更大模型。
  • 用户:在 CPU 上运行 Qwen3.5 等线性注意力模型时,可启用 --enable-chunked-prefill--enable-prefix-caching,显著提升吞吐并减少首 token 延迟(尤其长上下文场景)。
  • 系统:移除全局面上的限制后,所有 CPU 平台(无论是否 AMX)均尝试启用这些功能;AMX 平台之前被禁用,现在启用,可能增加 AMX 单元负载。
  • 团队:涉及 CPU 后端的多个模块(平台配置、模型运行器、工具函数、C++ 内核、测试),需确保后续线性注意力模型不回归。测试用例为后续开发提供了回归保障。
空指针风险 性能回退 测试覆盖有限

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论