Prhub

#47718 [ROCm][Perf] DSv4 two-stage compressor kernel for HCA prefill

原始 PR 作者 kliuae 合并时间 2026-07-16 07:31 文件变更 3 提交数 8 评论 3 代码增减 +528 / -0

执行摘要

两阶段 split compressor 优化 DSv4 prefill TTFT

在 ROCm 上,单阶段 compressor-norm-rope-cache fusion 以每 token 一个 workgroup 调度。当 compress ratio=128 时,只有位于 CR 边界的 token(每 128 个 token 一个)执行计算,导致大量 CU 闲置(例如 MI355 有 256 个 CU 但仅 64 个活跃)。为提升 occupancy,需将单个 program 拆分为更多独立 tasks,使 CU 饱和。

值得精读。此 PR 展示了如何通过 occupancy-targeted split 解决 kernel 级串行瓶颈,设计思路可推广到其他类似 pattern。关键设计决策:split 沿 head 维度无数据依赖、num_splits 动态计算、平台条件编译。

讨论亮点

tjtanaa 询问 split_decodes_and_prefills 是否应平台无关,kliuae 回应其实为 Triton 实现,理论上可跨平台,但目前仅 ROCm 经测试,故暂时限定 ROCm。

实现拆解

  1. 平台守卫与元数据扩展:在 compressor.py 中添加 _prefer_two_stage_compressor() 函数,仅在 ROCm 返回 True;在 CompressorMetadataBuilder.build() 中通过 split_decodes_and_prefills 获取 decode token 数,并新增 num_decode_tokens 字段。
  2. Scratch buffer 分配与前向分发:在 CompressorStateCache.__init__ 中,当 head_dim=512compress_ratio=128(无 overlap)时分配 fp32 格式的 _compress_scratch tensor [max_num_batched_tokens, head_dim];在前向方法中根据 num_decode_tokens 选择单阶段或两阶段调用。
  3. 两阶段 kernel 实现:在 fused_compress_quant_cache.py 中新增 _pick_compress_num_splits 根据 CU 数量和 boundary 数计算最优 splits(2 的幂,32 的倍数);_compress_gather_split_sparse_attn 为每个 (token, split) 启动一个 program,沿 head 维度分 tile 计算加权求和并写入 fp32 scratch;_finalize_norm_rope_quant_store_sparse_attn 在原有单 program 策略下完成 norm、rope、量化、插入。compress_norm_rope_store_two_stage_triton 对外暴露的接口。
  4. 测试覆盖:新增 test_fused_kv_insert_split 参数化 token 数(1/4/8/17)和 page 大小(16/64),验证两阶段 pipeline 输出与参考一致。
文件 模块 状态 重要度
vllm/models/deepseek_v4/common/ops/fused_compress_quant_cache.py 算子层 modified 7.69
vllm/models/deepseek_v4/compressor.py 压缩器 modified 7.48
tests/kernels/test_compressor_kv_cache.py 测试 modified 6.5

关键符号

_pick_compress_num_splits _compress_gather_split_sparse_attn _finalize_norm_rope_quant_store_sparse_attn _launch_two_stage_sparse_attn_compressor compress_norm_rope_store_two_stage_triton _prefer_two_stage_compressor test_fused_kv_insert_split

关键源码片段

vllm/models/deepseek_v4/compressor.py data-contract

修改压缩器入口,添加平台守卫和 scratch buffer 分配,根据 decode 数量选择两阶段路径。

def _prefer_two_stage_compressor() -> bool:
    """
    当前平台是否偏好两阶段 split compressor。
    Triton 实现在 AMD 上验证,因此暂时仅返回 ROCm。
    """
    return current_platform.is_rocm()# 在 CompressorStateCache.__init__ 中:
self._use_two_stage_fused_compressor = (
    _prefer_two_stage_compressor()
    and head_dim == 512
    and not self.overlap # cr>=128 才能开启
)
if self._use_two_stage_fused_compressor:
    self._compress_scratch = torch.empty(
        self.max_num_batched_tokens,
        self.head_dim,
        dtype=torch.float32,
        device=self.device,
    )# 在 forward 中:
elif self._use_two_stage_fused_compressor:
    compress_norm_rope_store_fn = compress_norm_rope_store_two_stage_triton
    extra_kwargs = {
        "num_decode_tokens": state_metadata.num_decode_tokens,
        "compress_scratch": self._compress_scratch,
    }
tests/kernels/test_compressor_kv_cache.py test-coverage

新增 Test F 覆盖两阶段 split compressor 的正确性,参数化 token 数和 page 大小。

@pytest.mark.skipif(
    not current_platform.is_rocm(),
    reason="two-stage split compressor is only enabled for ROCm at the moment",
)
@pytest.mark.parametrize("num_tokens", [1, 4, 8, 17])
@pytest.mark.parametrize("kv_block_size", [16, 64])
def test_fused_kv_insert_split(num_tokens: int, kv_block_size: int):
    """
    验证两阶段 split compressor 全流程正确性:
    state_cache gather → softmax 加权压缩 → RMSNorm → RoPE → FP8 量化 → paged insert。
    """
    HEAD_DIM = 512
    NOPE_DIM = 448
    ROPE_DIM = 64
    HEAD_BYTES = 584 # 448 fp8 + 128 bf16 + 8 uint8 scale
    RMS_EPS = 1e-6
    FP8_MAX = 448.0
    # ... 初始化数据,调用 _launch_two_stage_sparse_attn_compressor 后与参考比较

评论区精华

平台限定 vs 通用 设计

tjtanaa 评论 build 方法中无条件调用 split_decodes_and_prefills,可能应平台限定。kliuae 回复 Triton 实现可跨平台但当前仅 ROCm 测试,故暂时限定。

结论:作者使用 _prefer_two_stage_compressor 包裹来限定 ROCm。 · 已解决

风险与影响

新 kernel 仅作用于 ROCm 且 head_dim=512、cr>=128 场景,其他路径不受影响。主要风险:Triton 实现的跨平台兼容性未验证;scratch buffer 占用额外显存(max_batched_tokens * 512 * 4 bytes);两阶段分别执行增加了 kernel 启动开销,但在 prefill 下收益显著;解码路径仍使用单阶段。测试覆盖了多种 token 数与 page 大小,但未覆盖边缘 case(如 token 数极少)。

用户侧:ROCm 部署 DeepSeek-V4 模型时 prefill TTFT 降低 3-4%,吞吐小幅提升。解码路径无影响。系统侧:新增约 2MB/token(max_batched_tokens 1M+)scratch 空间。团队侧:需维护两套压缩路径,但差异化较小。对其他平台无影响。

仅 ROCm 验证 Triton 跨平台未测试 scratch 显存开销

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论