Prhub

#47896 [Kernel][ROCm][Perf] FlyDSL decode-attention kernel for 4-bit TurboQuant KV cache

原始 PR 作者 aditi-amd 合并时间 2026-08-11 23:45 文件变更 13 提交数 14 评论 47 代码增减 +6396 / -53

执行摘要

ROCm gfx950 新增 FlyDSL 4-bit TurboQuant 解码内核

PR 动机来自 AMD 团队在 agentic serving 场景对 4-bit TurboQuant KV-cache 解码性能的诉求:默认 Triton 解码在 gfx950 上未能充分利用 CDNA4 MFMA 指令与硬件转置,FlyDSL 提供的低层 DSL 可以写出单 wave CTA、LDS 常驻质心 LUT 等优化。PR body 同步附上 TurboQuant 算法背景博客(https://rocm.blogs.amd.com/artificial-intelligence/turboquant-vllm-agentic/README.html),说明其在 agentic vLLM 服务中的角色;设计上作为自动选路的替代方案(早期为 env 开关,最终改为 gfx950 + FlyDSL 可导入时的能力门控),确保默认路径保持不变。

值得精读。对 ROCm/AMD 内核开发者,这是 FlyDSL DSL 内核嵌入 vLLM attention 后端的完整样例;对平台维护者,_SegmBufPool + _max_capture_batch_size 解决多尺寸 CUDA-graph 捕获显存与地址稳定性的思路,以及 rocm.py backend 选择放宽的边界条件都值得借鉴。CUDA 用户无需深入内核细节,但可以关注后续 backend 抽象与配置字段的演进。

讨论亮点

评审主要围绕四个问题展开:一是 fxmarty-amd 质疑 rocm.py 中放宽 selected backend 校验会静默忽略无效的 --attention-backend(如 FLASHINFER_MLA_SPARSE_SM120),aditi-amd 将其收敛为仅 turboquant* KV dtype 下回退,其他 dtype 继续 fail loud;二是 fxmarty-amd 与 tjtanaa 都反对新增 VLLM_ROCM_TQ_FLYDSL_DECODE 等环境变量,tjtanaa 明确要求走 python 参数系统,最终提交删除了 env 开关改为 gfx950 能力自动选路,但 --attention-backend turboquant_flydsl 或配置字段留待后续;三是 fxmarty-amd 询问 FlyDSL 内核应托管在 vLLM 还是 AMD 专属仓库,aditi-amd 引用 PR #44400 作为先例,评审最终认可保留在 vLLM;四是 fxmarty-amd 建议以子类化替代 if/else 分支,aditi-amd 同意作为后续 RFC/PR。此外 tjtanaa 提出了测试平台导入保护与 AITER 版本策略问题,均已修复或回退。

实现拆解

  1. 新增 SoA 流水线基础:在 vllm/v1/attention/ops/turboquant_soa/ 下新增 triton_turboquant_store.py、triton_turboquant_decode.py、triton_turboquant_decode_v2.py 与 triton_turboquant_unified_attention.py。它们按 SoA 布局读写缓存——每个 block 的数据区与元数据区(k_norm/v_scale/v_zero)分离,目标是与 FlyDSL 内核共享同一套字节偏移约定;store 内核写入量化后的 KV,decode 内核在 tile 循环内直接反量化。
  2. 加入 FlyDSL 内核:新增 flydsl_kernels/tq_decode.py(GQA-8/16,Qwen 类)与 tq_decode_gqa6.py(GQA-6,MiniMax-M2.5),使用 FlyDSL DSL 生成 CDNA4 MFMA 指令;新增 flydsl_turboquant_decode.py 作为启动器,按形状分派、按 (Hk, num_partitions, max_blocks_per_seq, scale) 缓存编译产物,并通过 _SegmBufPool 用单一 max_B 桶复用中间输出,避免多尺寸 CUDA-graph 捕获带来的显存膨胀和每步 cudaMalloc。
  3. 接入 TurboQuantAttentionImpl:turboquant_attn.py 增加 _soa_imports 惰性导入、_use_flydsl/_soa_store 开关、_dispatch_decode_soa 分派,以及在 _ensure_on_device 中针对 FlyDSL 的 CUDA-graph 安全预热(预分配 arange/_cu_2、扩满 WorkspaceManager);_store_kv 与 continuation/prefill 路径在 SoA 开启时同步切换到 SoA Triton 实现,保证缓存布局与解码/续算路径一致。
  4. 调整 attention-backend 选择:rocm.py 的 get_attn_backend_cls 在 kv_cache_dtype 以 turboquant 开头时,允许显式选择的 backend 与该层不匹配并回退到自动选路,以支持边界层 AITER + turboquant 层 TURBOQUANT 的混合布局;其他 dtype 仍保持显式选择无效即报错。
  5. 测试与验证配套:新增 tests/kernels/turboquant/test_flydsl_turboquant_decode.py,用已知质心构造 SoA 4-bit KV cache,以纯 PyTorch fp32 attention 为 oracle,参数化覆盖 batch×seq_len×GQA×block_size,并且在非 gfx950 或 FlyDSL 不可导入时 module-level skip;PR body 附 GSM8K 精度与端到端性能数据。
文件 模块 状态 重要度
vllm/v1/attention/backends/turboquant_attn.py 注意力后端 modified 8.62
vllm/v1/attention/ops/flydsl_turboquant_decode.py 解码启动器 added 7.75
tests/kernels/turboquant/test_flydsl_turboquant_decode.py 单元测试 added 7.54
vllm/v1/attention/ops/flydsl_kernels/tq_decode.py 解码内核 added 7.75
vllm/v1/attention/ops/flydsl_kernels/tq_decode_gqa6.py 解码内核 added 7.75
vllm/v1/attention/ops/turboquant_soa/triton_turboquant_unified_attention.py Triton 内核 added 7.75
vllm/v1/attention/ops/turboquant_soa/triton_turboquant_decode.py Triton 解码 added 7.46
vllm/v1/attention/ops/turboquant_soa/triton_turboquant_decode_v2.py Triton 解码 added 7.46
vllm/v1/attention/ops/turboquant_soa/triton_turboquant_store.py KV 存储 added 7.46
vllm/platforms/rocm.py 平台层 modified 6.45
vllm/v1/attention/ops/turboquant_soa/__init__.py 包初始化 added 3.12
vllm/v1/attention/ops/flydsl_kernels/__init__.py 包初始化 added 3.02

关键符号

flydsl_turboquant_decode_attention is_flydsl_available is_flydsl_gqa6_available build_tq_decode_module build_tq_decode_gqa6_module triton_turboquant_store triton_turboquant_decode_attention_soa build_pair_lut _tq_load_k_tile _tq_fuse_q_rotation _soa_imports _dispatch_decode_soa _max_capture_batch_size _SegmBufPool.get _reference_attention test_flydsl_matches_reference

关键源码片段

vllm/v1/attention/ops/flydsl_turboquant_decode.py infrastructure

FlyDSL 启动器与运行时适配:负责 gfx950 能力探测、GQA-6 sibling 的 best-effort 导入、内核模块缓存,以及单桶 segm pool 对 CUDA-graph 多尺寸捕获的支持,是 FlyDSL 内核与 vLLM 解码循环之间的桥梁。

# vllm/v1/attention/ops/flydsl_turboquant_decode.py(新增)
# FlyDSL 能力探测:仅在 gfx950 且 FlyDSL 可导入时启用,否则回退 SoA Triton。_FLYDSL_AVAILABLE: bool | None = Nonedef is_flydsl_available() -> bool:
    """返回当前环境是否可用 FlyDSL TQ 解码(gfx950 + 可导入)。    GQA-6 sibling(tq_decode_gqa6)按 best-effort 导入:缺失时不影响
    Qwen 系 GQA-{8,16} 路径,只有 GQA-6 模型会回退到 SoA Triton。
    """
    if _FLYDSL_AVAILABLE is not None:
        return _FLYDSL_AVAILABLE
    try:
        from vllm.platforms.rocm import on_gfx950
        if not on_gfx950():
            _FLYDSL_AVAILABLE = False
            return False
        import flydsl.compiler as flyc
        import flydsl.expr as fx
        from flydsl._mlir import ir
        from flydsl.compiler.kernel_function import CompilationContext
        from flydsl.expr.typing import T
        from vllm.v1.attention.ops.flydsl_kernels import tq_decode as tq_mod
        _TQ_MOD = tq_mod
        _FLYDSL_AVAILABLE = True
        logger.info_once("FlyDSL TQ decode launcher: available")
    except Exception as ex: # noqa: BLE001
        _FLYDSL_AVAILABLE = False
        logger.warning_once(
            "FlyDSL TQ decode launcher: unavailable (%s). "
            "Falling back to SoA Triton decode.",
            ex,
        )
        return _FLYDSL_AVAILABLE
    # GQA-6 sibling 的导入独立于主路径,失败不影响 Qwen 支持。
    try:
        from vllm.v1.attention.ops.flydsl_kernels import tq_decode_gqa6 as tq_mod_gqa6
        _TQ_MOD_GQA6 = tq_mod_gqa6
        logger.info_once("FlyDSL TQ decode GQA-6 sibling: available (MiniMax-class)")
    except Exception as ex: # noqa: BLE001
        _TQ_MOD_GQA6 = None
        logger.info_once(
            "FlyDSL TQ decode GQA-6 sibling: not available (%s); "
            "GQA-6 models will fall back to SoA Triton decode.",
            ex,
        )
    return _FLYDSL_AVAILABLE
tests/kernels/turboquant/test_flydsl_turboquant_decode.py test-coverage

FlyDSL 解码内核的唯一正确性测试:构造 SoA 4-bit KV cache,并用纯 PyTorch fp32 attention oracle 校验输出,覆盖 batch×seq_len×GQA×block_size 组合,是非 gfx950 平台之外最重要的回归防线。

# tests/kernels/turboquant/test_flydsl_turboquant_decode.py(新增)
# 用 fp32 纯 PyTorch attention 作为 oracle,校验 FlyDSL 解码结果的正确性。def _reference_attention(q_bf16, k_ref, v_ref, seq_lens, scale):
    """对去量化后的 K/V 做 fp32 softmax attention (ground truth)。"""
    num_seqs, hq, d = q_bf16.shape
    hk = k_ref.shape[1]
    qg = hq // hk
    q = q_bf16.float().reshape(num_seqs, hk, qg, d)
    out = torch.zeros(num_seqs, hq, d, dtype=torch.float32)
    for s in range(num_seqs):
        for h in range(hk):
            ql = int(seq_lens[s].item())
            k = k_ref[s, h, :ql]
            v = v_ref[s, h, :ql]
            scores = (q[s, h] @ k.T) * scale
            m = scores.max(dim=-1, keepdim=True).values
            e = torch.exp(scores - m)
            p = e / e.sum(dim=-1, keepdim=True)
            out[s, h * qg : (h + 1) * qg] = p @ v
    return out@pytest.mark.parametrize("num_seqs", [1, 4, 16])
@pytest.mark.parametrize("seq_len", [256, 1024, 4096])
@pytest.mark.parametrize(
    "num_kv_heads,qg",
    [
        (8, 8), # Qwen2.5-72B class
        (8, 16), # Qwen3-32B class
        pytest.param(
            4, 6, # MiniMax-M2.5 class (GQA-6 sibling kernel)
            marks=pytest.mark.skipif(
                not is_flydsl_gqa6_available(),
                reason="FlyDSL GQA-6 sibling kernel not available",
            ),
        ),
    ],
)
@pytest.mark.parametrize("kv_block_size", [16, 32])
def test_flydsl_matches_reference(num_seqs, seq_len, num_kv_heads, qg, kv_block_size):
    """FlyDSL decode 必须与 fp32 attention (去量化 KV)一致。"""
    centroids, q_bf16, kv_cache, block_table, seq_lens, k_ref, v_ref = _build_cache(
        num_seqs, num_kv_heads, seq_len, qg, kv_block_size
    )
    scale = 1.0 / (HEAD_SIZE**0.5)
    identity = torch.eye(HEAD_SIZE, dtype=torch.float32, device="cuda")
    out = flydsl_turboquant_decode_attention(
        query=q_bf16,
        kv_cache=kv_cache,
        block_table=block_table,
        seq_lens=seq_lens,
        Pi=identity,
        centroids=centroids,
        scale=scale,
        mse_bits=4,
        key_packed_size=KEY_DATA_BYTES + 2,
        value_quant_bits=4,
        value_packed_size=KEY_DATA_BYTES + 4,
        key_fp8=False,
        norm_correction=False,
        PiT=identity.T.contiguous(),
        max_seq_len=seq_len,
        max_num_kv_splits=32,
        sinks=None,
    )
    ref = _reference_attention(q_bf16.cpu(), k_ref, v_ref, seq_lens.cpu(), scale)
    torch.testing.assert_close(out.cpu().float(), ref, atol=ATOL, rtol=0.0)

评论区精华

attention-backend 选择语义与 turboquant 混合后端 设计

fxmarty-amd 指出 vllm/platforms/rocm.py 中原先对所有无效 backend 直接 raise,而新改动会静默忽略 validate_configuration 的失败,例如指定 --attention-backend FLASHINFER_MLA_SPARSE_SM120 时不再报错。他同时质疑为何不要求用户对 turboquant_* KV 层显式传 --attention-backend TURBOQUANT,以及混合后端是否已有 AttentionConfig 抽象支撑。

结论:aditi-amd 将回退范围收敛到仅 turboquant* KV dtype,其他 dtype 下显式选择无效时仍 fail loud;混合后端是边界层跳过量化(--kv-cache-dtype-skip-layers)的既有设计,--attention-backend 目前是全局单值选择器。fxmarty-amd 认为该改动与 FlyDSL 集成关系不大,建议单独 PR 并评估跨平台一致性,此点未在本 PR 完全闭环。 · 部分解决

环境变量 vs 配置项 / 自动选路 设计

fxmarty-amd 问是否需要新增 VLLM_ROCM_TQ_FLYDSL_DECODE、VLLM_TQ_FLYDSL_WHT_BUTTERFLY 这类环境标志,并建议类似 nvfp4/mxfp4 的 backend 自动选择;tjtanaa 明确 'Let's not add flags for attention backend, all of those should go through python arguments system',并期望 --attention-backend turboquant_flydsl 或 --attention-config.turboquant-backend=flydsl 之类的配置。

结论:最终提交删除了环境变量,改为 gfx950 + FlyDSL 可导入时的能力自动选路(on_gfx950() gating);但 --attention-backend turboquant_flydsl 或配置字段的落地方案留待后续 PR/RFC。 · 部分解决

FlyDSL 内核托管位置 设计

fxmarty-amd 问 AMD 专用 FlyDSL 内核是放在 vLLM 主仓还是 aiter/flydsl/amd-quark 等 AMD 仓库,vLLM 只做前端;@mgoin 被征询意见。

结论:aditi-amd 以 PR #44400 为先例将内核随 vLLM 维护;fxmarty-amd 转述 'keeping AMD-only kernels in vllm-project/vllm is fine',话题关闭。 · 已解决

子类化改造 if/else 分支 设计

fxmarty-amd 建议 TurboQuantAttentionImpl 保持默认 Triton 路径不变,新建 FlyDSLTurboQuantAttentionImpl 子类覆盖 store/continuation/dequant/decode,避免在基类里堆 if/else;aditi-amd 认可,计划在 env/ 配置决策定稿后以独立 RFC/PR 落地。

结论:本 PR 仍以内联分支形式合入,子类化重构作为后续项挂起。 · 未解决

测试的非 ROCm 导入保护 测试

tjtanaa 指出 tests/kernels/turboquant/test_flydsl_turboquant_decode.py 在模块顶层导入 vllm.platforms.rocm 会在非 ROCm 平台失败,需用 current_platform.is_rocm() 守卫;aditi-amd 确认修复。

结论:测试改为 module-level skip,先判断 current_platform.is_rocm() 和 on_gfx950(),再在条件内导入 ROCm/FlyDSL 依赖,CI 非 ROCm 平台可干净跳过。 · 已解决

AITER 版本兼容与 lazy import 回退 other

fxmarty-amd 与 tjtanaa 指出 rocm_aiter_moe.py 和 _aiter_ops.py 里为兼容旧版 aiter 的惰性导入 / 能力探测与 vLLM 的 AITER 版本策略冲突(只支持 Dockerfile 中的特定版本,不做前向 / 反向兼容),要求 revert。

结论:aditi-amd 接受并全部 revert,使用上游 Docker 镜像自带 AITER。 · 已解决

v4 命名与 Triton 回退语义 style

fxmarty-amd 问测试文件名中的 v4 来源;aditi-amd 解释是内部迭代编号,会删除。另一条线上,fxmarty-amd 觉得 VLLM_ROCM_TQ_FLYDSL_DECODE=1 时仍可能落到 Triton 内核令人困惑,建议只保留最优实现或改更贴切的名字;aditi-amd 解释 FlyDSL 只覆盖 gfx950 特定 profile,Triton SoA 是通用回退,两者都必须保留。

结论:v4 字样全部移除;回退命名语义未在本 PR 修改,作者表示可后续探讨。 · 部分解决

风险与影响

  1. 平台强绑定:FlyDSL 内核仅面向 gfx950/CDNA4,其他 ROCm 卡依赖 Triton SoA 回退;若回退路径未覆盖的边界形状(如 key_fp8=1、非 4-bit、HEAD_SIZE≠128)被错误路由,性能收益消失但功能不受影响。
  2. 布局混用风险:同一 kv_cache 张量只能有一种字节约定,turboquant_attn.py 中 SoA 开启时 _store_kv 与 continuation 路径必须同步切换,否则 AoS 解码读 SoA 缓存会拿到错误的 k_norm/v_scale/v_zero 偏移,导致多轮或前缀缓存请求精度崩溃。
  3. attention-backend 选择行为变化:rocm.py 对 turboquant* dtype 放宽校验,虽限定前缀且其余 dtype 仍报错,但跨平台一致性与用户对显式 --attention-backend 的预期仍需澄清。
  4. CUDA-graph 捕获路径:新增的预热、segm pool 与 arange cache 依赖 max_model_len 与 capture 配置;超出最大捕获 batch 时回退 eager,已通过 _max_capture_batch_size 取 max 缓解。
  5. 测试覆盖有限:核心正确性测试只在 ROCm/gfx950 且 FlyDSL 可用时运行,CUDA 与多数 ROCm CI 均为 skip,回归面有限。

影响范围集中:仅当使用 4-bit TurboQuant KV(如 --kv-cache-dtype turboquant_4bit_nc)且处于 ROCm gfx950 + FlyDSL 可导入时触发,默认路径完全不变。对 MI355X 用户,agentic 或长上下文 decode 可获得约 4.1× 相对默认 Triton 解码的吞吐提升,精度整体中性(GSM8K 差距 <0.4pp);对其他硬件用户无行为变化。系统层面,attention-backend 选择逻辑新增 turboquant 前缀分支,混合后端(边界层 AITER + 主体 TURBOQUANT)成为受支持组合。团队层面,AMD 侧需要持续维护 FlyDSL DSL 内核与 vLLM 版本的兼容性,并跟进 --attention-backend turboquant_flydsl、TurboQuantAttentionImpl 子类化等后续抽象演进。

gfx950-only 内核 attention-backend 选择行为变更 CUDA graph 捕获路径改动 测试覆盖仅限 gfx950 SoA/AoS 布局混用风险 新增外部 DSL 依赖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论