Prhub

#48018 [Kernel] ReplaySSM: cache SSM inputs for faster Mamba2 standard decode

原始 PR 作者 Johnny-Liou 合并时间 2026-07-25 00:39 文件变更 23 提交数 30 评论 45 代码增减 +2954 / -31

执行摘要

ReplaySSM 缓存 SSM 输入加速 Mamba2 标准解码

SSM 解码内存受限,每步加载和存储 (nheads, d, n) 循环状态,存储开销主导 HBM 流量。根据 PR body 中的带宽分析表,在批大小为 512 时基线状态更新内核仅达到 84.6% 的峰值 DRAM 带宽,表明 HBM 访问是瓶颈。ReplaySSM 维护一个小的输入环缓冲和周期检查点,大多数步骤直接从检查点加缓冲读取输出,而无需物化完整状态,每 L 步才写回一次完整状态,从而减半存储流量。

值得仔细阅读,尤其关注其分阶段实现策略、环缓冲+周期检查点的设计模式、以及 CUDA graphs 下的每行动态分支处理。Review 中关于随机舍入和请求驱逐的讨论展示了重要的边界情况处理。建议相关开发人员了解 ReplaySSM 的配置约束,为后续扩展做准备。

讨论亮点

Review 中主要讨论点:

  • 随机舍入支持:tomeraas91 要求添加对量化模型(如 Nemotron-3 NVFP4)的随机舍入支持,因为基线依赖 SR 保证正确性。Johnny-Liou 在后续提交中实现了 SR,但指出当前 Triton 实现延迟较高,后续将推出更高效的 Blackwell 特化内核(CuteDSL)。双方同意当前实现已能避免精度问题,未来优化可跟进。

  • 元数据构建器测试缺失:tomeraas91 指出最危险的新逻辑是 flush 检测和单 token 预填充边界情况,要求添加类似 test_gdn_metadata_builder.py 的测试。Johnny-Liou 随后添加 test_replayssm_metadata_builder.py,全面覆盖各种场景(fresh/resume/align/flush 边界)。

  • 请求驱逐 bug 修复:初始实现假设解码开始时检查点位于 prompt 边界,但抢占后此假设不成立。tomeraas91 发现此问题,Johnny-Liou 引入 replayssm_decode_base 字段在每次恢复时重锚定写位置,并添加了 preemption 测试。

  • 前缀缓存支持:最初要求 mamba_cache_mode=none,但 tomeras91 询问是否可支持前缀缓存。Johnny-Liou 在后续提交中增加了 align 模式支持,并验证了前缀缓存命中。

  • 冗余检查与代码风格:tomeraas91 指出 mixer2 中多处运行时检查是死代码(因上游已保证),Johnny-Liou 删除之。同时统一了配置属性命名。

  • 模型泛化:tomeraas91 质疑为何仅限 Nemotron-H 启用。Johnny-Liou 回答当前仅在该模型上验证过,未来可通过预调优配置轻松扩展至其他 Mamba2 模型。

实现拆解

  1. 新增 Triton 内核:在 vllm/model_executor/layers/mamba/ops/selective_state_update_replayssm_output_only.py 中实现 selective_state_update_replayssm_output_only 函数,包含两个 Triton 内核:_replayssm_output_only_precompute_kernel(处理非刷新行的预计算)和 _replayssm_output_only_kernel(根据 is_flush 每行分支选择刷新或快照路径)。

  2. 集成到 MambaMixer2:在 mamba_mixer2.py 中根据 cache_config.use_replayssm 标志初始化 ReplaySSM 相关属性,扩展 kv_cache 元组为 5 个张量(conv_state, ssm_state, x_cache, dt_cache, B_cache),并在 conv_ssm_forward 中根据标志派发至 ReplaySSM 解码路径。

  3. 元数据构建器扩展:在 mamba_attn.pyBaseMambaAttentionMetadataBuilder 中新增 decode_write_pos_ddecode_is_flush_d 张量分配,并在 builder 的 build 方法中基于 replayssm_decode_base 计算每行的 write_posis_flush,同时支持 align 模式的重锚定。

  4. 配置层验证:在 vllm/config/vllm.pyvalidate_mamba_cached_kernel 中确保 ReplaySSM 不能与推测解码、KV 传输/分离式服务及除 none/align 外的缓存模式同时使用。

  5. 测试覆盖:新增单元测试覆盖 Triton 内核的正确性(test_replayssm_standard_decode_mamba2.py)、元数据构建器(test_replayssm_metadata_builder.py)、端到端 logprobs 一致性(test_replayssm_decode.py,含 TP2 和前缀缓存 align 模式)、以及预填充-解码等价性(test_replayssm_prefill_decode_equivalence_mamba2.py)。

文件 模块 状态 重要度
vllm/model_executor/layers/mamba/mamba_mixer2.py Mamba2 层 modified 8.25
vllm/model_executor/layers/mamba/ops/selective_state_update_replayssm_output_only.py SSM 内核 added 7.93
vllm/v1/attention/backends/mamba_attn.py 注意力后端 modified 7.44
vllm/config/vllm.py 配置 modified 7.05
tests/kernels/mamba/test_replayssm_standard_decode_mamba2.py 内核测试 added 8.15

关键符号

selective_state_update_replayssm_output_only _replayssm_output_only_precompute_kernel _replayssm_output_only_kernel append_replayssm_ring supports_replayssm validate_mamba_cached_kernel get_replayssm_config

关键源码片段

vllm/model_executor/layers/mamba/mamba_mixer2.py core-logic

Mamba2 混合器的核心修改,集成 ReplaySSM 解码路径,扩展 kv_cache 为 5 个张量,并根据配置选择内核。

# vllm/model_executor/layers/mamba/mamba_mixer2.py ( 部分初始化代码 )# 导入新增的 ReplaySSM 内核
from vllm.model_executor.layers.mamba.ops.selective_state_update_replayssm_output_only import (
    selective_state_update_replayssm_output_only,
)class MambaMixer2(nn.Module):
    def __init__(self, ...):
        # ... 原有初始化 ...
​
        # 从 cache_config 读取 ReplaySSM 标志和缓冲长度
        self.use_replayssm = (
            cache_config.use_replayssm if cache_config is not None else False
        )
        self.replayssm_buffer_len = (
            cache_config.replayssm_buffer_len
            if cache_config is not None and cache_config.use_replayssm
            else None
        )
​
        # TP 头部数必须整除,否则无法分片
        if self.use_replayssm and self.num_heads % self.tp_size != 0:
            raise ValueError(
                "--use-replayssm requires tensor-parallel heads to divide evenly"
            )
​
        # kv_cache 元组扩展:基础状态 (conv_state, ssm_state) 或
        # ReplaySSM 状态 (conv_state, ssm_state, x_cache, dt_cache, B_cache)
        _n_state = 5 if self.use_replayssm else 2
        self.kv_cache = tuple(torch.tensor([]) for _ in range(_n_state))
​
        # ... 后续初始化 ...
​
    def conv_ssm_forward(self, ...):
        # 根据 use_replayssm 解包额外的缓存
        if self.use_replayssm:
            x_cache, dt_cache, B_cache = self.kv_cache[2:]
        else:
            x_cache = dt_cache = B_cache = None
​
        # ... 中间计算 ...
​
        # 如果启用 ReplaySSM,调用 output_only 内核
        if self.use_replayssm:
            selective_state_update_replayssm_output_only(
                ssm_state,
                hidden_states_d, dt_d, A_d, B_d, C_d, D_d,
                x_cache, dt_cache, B_cache,
                write_pos=attn_metadata.decode_write_pos_d,
                is_flush=attn_metadata.decode_is_flush_d,
                state_batch_indices=...,
                out=preallocated_ssm_out_d,
                dt_bias=dt_bias, dt_softplus=True,
            )
        else:
            # 原 baseline 的 selective_state_update
            selective_state_update(
                ssm_state, hidden_states_d, dt_d, A_d, B_d, C_d, D_d,
                dt_bias=dt_bias, dt_softplus=True,
                state_batch_indices=...,
                out=preallocated_ssm_out_d,
            )

评论区精华

随机舍入支持 设计

tomeraas91 要求添加随机舍入支持以用于量化模型,Johnny-Liou 在后续提交中实现,但指出当前 Triton 实现延迟较高,后续将推出 Blackwell 特化内核。

结论:已添加 SR 支持,当前实现可避免精度问题;未来通过 CuteDSL 内核优化性能。 · 已解决

元数据构建器测试缺失 测试

tomeraas91 指出最危险的新逻辑是 flush 检测,要求添加类似 GDN 的测试。Johnny-Liou 随后添加 test_replayssm_metadata_builder.py,覆盖各种场景。

结论:已添加全面测试,包括 fresh/resume/align/flush 边界。 · 已解决

请求驱逐 bug 修复 正确性

tomeraas91 发现 write_pos 计算在请求被抢占后失效,因为 checkpoint 不再位于 prompt 边界。Johnny-Liou 引入 replayssm_decode_base 修复,并添加 preemption 测试。

结论:已修复,通过重锚定 base 在每次恢复时更新。 · 已解决

前缀缓存与 align 模式支持 设计

最初 ReplaySSM 要求 mamba_cache_mode=none,tomeraas91 询问是否可支持前缀缓存。Johnny-Liou 在后续提交中增加了 align 模式支持。

结论:已支持 align 模式,并验证前缀缓存命中。 · 已解决

仅限 Nemotron-H 模型启用 设计

tomeraas91 质疑为何 ReplaySSM 仅对 Nemotron-H 启用,认为应是 Mamba2 块属性而非模型属性。Johnny-Liou 同意未来可扩展,但当前仅在该模型上验证过,需配合预调优配置。

结论:当前仅限 Nemotron-H,后续可通过添加 suppoerts_replayssm 标志和预调优配置轻松扩展。 · 已解决

风险与影响

  1. 兼容性风险:ReplaySSM 与推测解码、KV 传输/分离式服务、all 前缀缓存模式不兼容,已在配置层显式抛出错误,但若用户误用可能困惑。
  2. 数值精度风险:在 bf16 状态模式下,ReplaySSM 与基线 kernel 输出不同(因基线每步执行 fp32->bf16 降精度,而 ReplaySSM 以 fp32 累积缓冲后一次写回)。虽在测试中验证了正确性,但用户需注意 logprobs 可能细微差异。
  3. 性能风险:小 batch(1-8)下 ReplaySSM 有轻微减速(0.96x-1.00x),仅在大 batch 下才有收益。用户需根据负载决定是否启用。
  4. 维护风险:新增 Triton 内核和集成代码增加了 Mamba2 路径的复杂性,未来 Mamba1 或其他 SSM 变体可能需要类似适配。
  5. 测试覆盖不足:目前仅针对 Nemotron-3 系列模型测试,其他 Mamba2 模型(如 Hymba、Phi-Mamba)未经验证,可能需要调优内核配置。

用户影响:使用 Mamba2 混合模型(如 Nemotron-3)的用户可通过启用 --use-replayssm 获得服务批处理性能提升(大 batch 端到端 1.1x-1.86x)。需注意禁用前缀缓存(none 模式)或改用 align 模式。系统影响:每请求增加约 3 * buffer_len * nheads * headdim 字节的额外显存用于环缓冲,但默认 buffer_len=16,开销可忽略。团队影响:该 PR 是完整 ReplaySSM 设计的第一阶段,后续将扩展至 Gated DeltaNet 和推测解码,影响面将扩大。

随机舍入支持延迟高 模型特定仅限 Nemotron-H 小 batch 性能退化 与推测解码不兼容 数值精度差异在 bf16 状态

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论