Prhub

#43083 Tuning script and configs for Triton Mamba SSU kernel

原始 PR 作者 danisereb 合并时间 2026-05-25 01:12 文件变更 13 提交数 24 评论 28 代码增减 +1770 / -91

执行摘要

为 Mamba SSU 内核添加自动调优框架与预配置,性能提升最高 1.6x

Issue #33034 提出需要为 selective_state_update 内核提供与 fused MoE 类似的调优能力,通过 JSON 配置文件实现设备特定的最佳参数选择,从而在不同 GPU 上获得一致的高性能,并降低对新设备硬编码的需求。PR #41398 的前期尝试需要进一步完善和通用化。

该 PR 值得精读,尤其是如何将硬编码调优参数优雅地替换为数据驱动的配置系统,包括配置生成、缓存、环境变量覆盖、fallback 和测试策略。设计模式可复用至其他 kernel 的调优。

讨论亮点
  • 配置目录位置争议:tomeras91 建议将 configs 放在 ops/ 子目录下而非 mamba/ 层,以对齐 MoE 模式并避免与其他 kernel 混淆。作者采纳并调整。
  • BF16 配置复用:tomeras91 指出 BF16 和 FP16 对 kernel 表现相同,应共享配置。作者通过 _canonical_cache_dtype 映射实现。
  • 缓存一致性:tomeras91 发现 @functools.cache@functools.lru_cache 混用,作者统一为 @functools.cache
  • 调优参数拆分:tomeras91 建议将 --effective-batches 拆分为 --batch-sizes--nheads,使脚本更自然,作者照做。
  • 测量重复问题:tomeras91 指出对比表中 Tuned[us] 和 Heur[us] 数值差异源于独立两轮测量,作者改为重用测量结果。
  • JSON 安全解析:gemini-code-assist[bot] 建议用 k.isdigit() 过滤非整数键(如 triton_version),防止用户自定义配置文件崩溃,作者采纳。

实现拆解

  1. 配置加载基础设施:在 vllm/model_executor/layers/mamba/ops/mamba_ssm.py 中新增 get_ssm_config_file_nameget_ssm_configs_try_get_optimal_ssm_config_cached 等函数。通过 @functools.cache 缓存配置加载结果,使用 _canonical_cache_dtypebfloat16 映射到 float16 以复用配置。新增 override_ssm_config 上下文管理器用于测试或动态覆盖。

  2. 调优脚本:新建 benchmarks/kernels/benchmark_selective_state_update.py,仿照 fused MoE 调优流程。定义搜索空间(BLOCK_SIZE_M 上限为 next_pow2(headdim),num_warps 从 1 到 8),使用 CUDA graph 捕获减少延迟测量抖动。通过 --batch-sizes--nheads 参数生成 effective_batch 网格,支持 --save-configs 输出 JSON、--validate 验证正确性、--compare 与默认启发式对比。

  3. 内核修改selective_state_update 函数删除原有基于 dstate 的 if-else 硬编码,改调 try_get_optimal_ssm_config 获取配置。若配置不存在或加载失败则回退到 _get_default_ssm_launch_config 启发式,确保兼容。保留 is_blackwell 标志控制部分特性。

  4. 测试与重构:新增 tests/kernels/mamba/test_mamba_ssm_configs.py 覆盖文件名格式、VLLM_TUNED_CONFIG_FOLDER 环境变量、回退逻辑、effective_batch 插值、非 dict JSON 和空配置等边界情况。将 selective_state_update_reftest_mamba_ssm.py 迁移到 tests/kernels/mamba/utils.py 避免重复,并在测试和基准脚本中统一引用。

  5. 预配置 JSON 文件:在 vllm/model_executor/layers/mamba/ops/configs/selective_state_update/ 下生成 8 个配置文件(涵盖 headdim=64, dstate=128,float16/float32 缓存类型,设备包括 NVIDIA_B200, NVIDIA_GB200, NVIDIA_H100_80GB_HBM3),每个文件包含 effective_batch 从 8 到 262144 的 (BLOCK_SIZE_M, num_warps) 配置。

配套修改 vllm/envs.py 添加 VLLM_TUNED_CONFIG_FOLDER 环境变量,允许用户指定自定义配置目录。

文件 模块 状态 重要度
benchmarks/kernels/benchmark_selective_state_update.py 调优工具 added 8.98
vllm/model_executor/layers/mamba/ops/mamba_ssm.py SSM 内核 modified 7.64
tests/kernels/mamba/test_mamba_ssm_configs.py 配置测试 added 7.42
tests/kernels/mamba/utils.py 测试工具 added 5.99
tests/kernels/mamba/test_mamba_ssm.py SSM 测试 modified 5.33
vllm/model_executor/layers/mamba/ops/configs/selective_state_update/headdim=64,dstate=128,device_name=NVIDIA_B200,cache_dtype=float32.json 预配置 added 4.89

关键符号

get_ssm_configs try_get_optimal_ssm_config benchmark_config _block_size_m_choices tune_dstate validate_configs save_configs override_ssm_config

关键源码片段

benchmarks/kernels/benchmark_selective_state_update.py benchmark-tool

新增调优脚本,是 PR 核心产出之一,用于自动搜索最优 BLOCK_SIZE_M 和 num_warps 并生成 JSON 配置。

# 调优脚本核心:测试单个 (BLOCK_SIZE_M, num_warps) 配置的延迟
def benchmark_config(
    batch: int, nheads: int, dim: int, dstate: int, ngroups: int,
    block_size_m: int, num_warps_val: int, dtype: torch.dtype,
    state_dtype: torch.dtype | None = None,
    num_iters: int = 100, num_warmup: int = 20, graph_batch_size: int = 10,
) -> float | None:
    # 生成随机输入,形状匹配 kernel 预期
    state, x, dt, A, B, C, D, dt_bias, out = _make_inputs(
        batch, nheads, dim, dstate, ngroups, dtype, state_dtype=state_dtype
    )
    # 包装内核调用,传入待测的 BLOCK_SIZE_M 和 num_warps
    def _call_kernel() -> None:
        selective_state_update(
            state, x, dt, A, B, C, D=D, z=None, dt_bias=dt_bias,
            block_size_m=block_size_m, num_warps_val=num_warps_val,
        )
    # 使用 CUDA graph 捕获并回放,消除 Python 调度开销
    # (内部实现包括预热、graph 捕获、多次回放取中位数)
    elapsed_us = ... # 省略具体实现
    return elapsed_us
vllm/model_executor/layers/mamba/ops/mamba_ssm.py core-logic

核心文件,添加 JSON 配置加载、缓存、设备名获取、规范化 dtype 等基础设施,并修改 selective_state_update 内核动态加载配置。

# 加载 JSON 配置文件并返回 effective_batch -> (BLOCK_SIZE_M, num_warps) 映射
def get_ssm_configs(
    headdim: int, dstate: int, cache_dtype: str
) -> dict[int, Any] | None:
    # 将 bfloat16 规范化到 float16(共用配置)
    cache_dtype = _canonical_cache_dtype(cache_dtype)
    device_name = get_ssm_device_name()
    json_file_name = get_ssm_config_file_name(
        headdim, dstate, cache_dtype, device_name
    )
    # 优先从 VLLM_TUNED_CONFIG_FOLDER 环境变量指定的目录加载
    config_file_paths: list[str] = []
    user_defined_config_folder = envs.VLLM_TUNED_CONFIG_FOLDER
    if user_defined_config_folder is not None:
        config_file_paths.append(
            os.path.join(user_defined_config_folder, json_file_name)
        )
    # 后备使用仓库自带的配置
    config_file_paths.append(os.path.join(_CONFIGS_DIR, json_file_name))
​
    for path in config_file_paths:
        if os.path.exists(path):
            with open(path) as f:
                logger.info_once(
                    "Using SSM config from %s for selective_state_update", path
                )
                raw = json.load(f)
                # 弹出辅助字段 triton_version,只保留整数字符串键
                raw.pop("triton_version", None)
                # 通过 k.isdigit() 过滤非数字键(防止用户自定义配置注入额外字段)
                return {int(k): v for k, v in raw.items() if k.isdigit()}
    # 无配置文件时返回 None,由调用方回退到启发式
    return None
tests/kernels/mamba/test_mamba_ssm_configs.py test-coverage

新增单元测试,全面覆盖配置加载、环境变量覆盖、回退、插值及边界情况,保障可靠性。

# 验证 VLLM_TUNED_CONFIG_FOLDER 环境变量优先于内置配置
def test_env_override_loads_custom_config(monkeypatch, tmp_path):
    # 在临时目录写入自定义配置
    _write_config(
        tmp_path, dstate=16,
        payload={"1": {"BLOCK_SIZE_M": 4, "num_warps": 1}},
    )
    # 设置环境变量指向临时目录
    monkeypatch.setenv("VLLM_TUNED_CONFIG_FOLDER", str(tmp_path))
    _clear_caches() # 清除函数级缓存,确保重新加载
    cfg = get_ssm_configs(_HEADDIM, 16, _CACHE_DTYPE)
    assert cfg is not None
    # 确认加载到的配置与写入的自定义配置一致
    assert cfg[1] == {"BLOCK_SIZE_M": 4, "num_warps": 1}
    _clear_caches()

评论区精华

配置目录位置 设计

tomeras91 建议将 configs 置于 ops/ 下而非 mamba/ 层,以对齐 MoE 模式并避免与其他 kernel 冲突。

结论:作者接受并移动 configs 目录到 selective_state_update/。 · 已解决

BF16 配置复用 正确性

tomeras91 指出 BF16 和 FP16 对 kernel 行为相同,应共享配置。

结论:通过 _canonical_cache_dtype 将 bfloat16 映射到 float16。 · 已解决

缓存装饰器一致性 style

tomeras91 指出 get_ssm_configs 使用 lru_cache 而另一函数使用 cache。

结论:统一为 @functools.cache。 · 已解决

effective_batch 参数化 设计

tomeras91 建议拆分 --effective-batches 为 --batch-sizes 和 --nheads。

结论:作者拆分参数,提升脚本可用性。 · 已解决

测量重复问题 other

tomeras91 发现对比表中 Tuned[us] 和 Heur[us] 数值来自独立两轮测量,导致非严格公平对比。

结论:作者修改为重用测量结果。 · 已解决

JSON 安全解析 安全

gemini-code-assist[bot] 建议用 k.isdigit() 过滤非整数键,防止用户自定义配置字段导致崩溃。

结论:作者采纳,增加过滤。 · 已解决

风险与影响

  1. 性能退化风险:若配置加载失败或有效 effective_batch 外推不佳,内核回退到启发式。但回退路径已验证且与旧行为一致,风险较低。
  2. 模型兼容性:部分模型使用 BF16 SSM cache(如 Granite 4H、Codestral Mamba 7B),配置映射到 FP16 后未单独调优,但理论性能一致,无正确性风险。
  3. 新增环境变量VLLM_TUNED_CONFIG_FOLDER 如果指向不存在或格式错误的文件,会影响内核初始化,但代码已做 fallback 并记录日志。
  • 用户:使用 Mamba 模型(如 Nemotron)的用户将自动获得性能提升,无需手动干预;高级用户可通过环境变量使用自定义配置。
  • 系统:预配置 JSON 随仓库发布,不会增加安装负担。基准脚本可用于新 GPU 的调优。
  • 团队:建立了类似 fused MoE 的标准化调优流程,降低后续维护和扩展成本。
配置加载失败回退路径验证 BF16 用户可能忽略手动调优 新增环境变量兼容性

关联 Issue

#33034 [Feature][Help Wanted]: Add tuning script and config files for Mamba selective_state_update kernel

完整报告

参与讨论