Prhub

#45103 [ROCm][DSV4][Perf] Fuse inverse-RoPE and cache bf16 wo_a in o-projection

原始 PR 作者 Fangzhou-Ai 合并时间 2026-06-13 04:57 文件变更 2 提交数 4 评论 9 代码增减 +341 / -59

执行摘要

融合逆 RoPE 与 wo_a 缓存,ROCm DSV4 decode 性能提升 9-16%

从PR body可知,ROCm上DSV4的o-projection(rocm_inv_rope_einsum)逆RoPE由约10个小PyTorch kernel组成(clone, index_select, repeat_interleave, neg, stack, cat, cast),且静态fp8 wo_a权重每次步骤都需反量化。这些操作在profile中显示为最大的拷贝/乘法kernel,成为decode性能瓶颈。

值得精读,该PR展示了如何通过Triton kernel融合和权重缓存实现显著性能优化,设计思路清晰,测试充分。尤其是将多步小kernel合并为单launch、以及惰性缓存模式的实践,对其他算子优化有参考价值。

讨论亮点

核心review讨论集中于:

  • tjtanaa 提出_can_use_fused_inv_rope始终为True,建议移除回退路径并直接使用融合kernel。作者响应“helper removed”,删除了相关函数和回退代码。
  • dllehr-amd 询问测试是否应与官方DeepseekV4ScalingRotaryEmbedding对比,作者确认已更新测试,直接使用官方旋转嵌入的forward_native作为参考。
  • dllehr-amd 最终批准。

实现拆解

  1. 引入Triton融合kernel:在rocm_aiter_mla_sparse.py中定义_inverse_rope_gptj_kernel(@triton.jit),一次launch完成逆RoPE,直接输出bf16。该kernel假设非neox布局、rope_dim为偶数、连续张量。
  2. 替换逆RoPE调用:将rocm_inv_rope_einsum中的_apply_inv_rope_ref(包含PyTorch回退)替换为_fused_inverse_rope_gptj调用,并根据review建议移除了_can_use_fused_inv_rope函数,使融合路径成为唯一路径(通过assert确保条件满足)。
  3. 缓存反量化wo_a权重:在_get_cached_wo_a_bf16中,首次调用时将fp8权重反量化并存储为模块属性wo_a._dsv4_wo_a_bf16,后续步骤直接返回缓存,避免每步重复反量化。
  4. 添加配套测试:在test_rocm_triton_attn_dsv4.py中增加三个测试类:与官方DeepseekV4ScalingRotaryEmbedding.forward_native对比的test_fused_inverse_rope_gptj_matches_rotary_native(参数化token数、头数、位置dtype)、空张量测试test_fused_inverse_rope_gptj_empty、以及缓存功能测试test_get_cached_wo_a_bf16_plain_caches。
文件 模块 状态 重要度
vllm/v1/attention/ops/rocm_aiter_mla_sparse.py 注意力层 modified 7.41
tests/kernels/attention/test_rocm_triton_attn_dsv4.py 测试 modified 7.19

关键符号

_inverse_rope_gptj_kernel _fused_inverse_rope_gptj _get_cached_wo_a_bf16 rocm_inv_rope_einsum

关键源码片段

vllm/v1/attention/ops/rocm_aiter_mla_sparse.py core-logic

核心优化实现文件,包含 Triton 融合 kernel、逆 RoPE 替换、wo_a 缓存逻辑。

# vllm/v1/attention/ops/rocm_aiter_mla_sparse.py ( 片段 )@triton.jit
def _inverse_rope_gptj_kernel(
    o_ptr, # [T, H, D] 输入张量指针
    out_ptr, # [T, H, D] 输出张量指针 (bf16)
    pos_ptr, # [T] 位置索引指针
    cos_sin_ptr, # [P, rope_dim] cos/sin 缓存 (fp32, cos 在前半 , sin 在后半 )
    s_t, s_h, # 输入行步长 ( 最后一维连续 )
    os_t, os_h, # 输出行步长
    cs_stride, # cos_sin_cache 行步长
    NOPE: tl.constexpr, # 非 rope 维度数 ( 直接透传 )
    HALF: tl.constexpr, # rope_dim // 2
    BLOCK_NOPE: tl.constexpr,
    BLOCK_HALF: tl.constexpr,
):
    '''融合逆 GPT-J RoPE 在末尾 rope_dim 维度。    实现与 DeepseekV4ScalingRotaryEmbedding.forward_native(inverse=True)
    相同的计算,但直接以 bf16 写出,替代了约 10 个小 kernel 的链式操作。
    '''
    t = tl.program_id(0)
    h = tl.program_id(1)
    in_base = t * s_t + h * s_h
    out_base = t * os_t + h * os_h
​
    # NoPE 部分直接透传 ( 转换为 bf16)
    n = tl.arange(0, BLOCK_NOPE)
    nmask = n < NOPE
    vals = tl.load(o_ptr + in_base + n, mask=nmask)
    tl.store(out_ptr + out_base + n, vals.to(tl.bfloat16), mask=nmask)
​
    # RoPE 部分 : out_even = a*cos + b*sin, out_odd = b*cos - a*sin
    # (a = 偶序号输入 , b = 奇序号输入 ; sin 取负以实现逆旋转 )
    pos = tl.load(pos_ptr + t).to(tl.int64)
    k = tl.arange(0, BLOCK_HALF)
    kmask = k < HALF
    a = tl.load(o_ptr + in_base + NOPE + 2 * k, mask=kmask).to(tl.float32)
    b = tl.load(o_ptr + in_base + NOPE + 2 * k + 1, mask=kmask).to(tl.float32)
    cos = tl.load(cos_sin_ptr + pos * cs_stride + k, mask=kmask)
    sin = tl.load(cos_sin_ptr + pos * cs_stride + HALF + k, mask=kmask)
    out_even = a * cos + b * sin
    out_odd = b * cos - a * sin
    tl.store(out_ptr + out_base + NOPE + 2 * k, out_even.to(tl.bfloat16), mask=kmask)
    tl.store(out_ptr + out_base + NOPE + 2 * k + 1, out_odd.to(tl.bfloat16), mask=kmask)
​
​
def _get_cached_wo_a_bf16(
    wo_a: torch.nn.Module,
    hidden_dim: int,
    o_lora_rank: int,
    n_local_groups: int,
) -> torch.Tensor:
    '''惰性缓存反量化后的 bf16 wo_a 权重。首次调用时缓存,后续直接返回。'''
    if not hasattr(wo_a, '_dsv4_wo_a_bf16'):
        # 首次:执行反量化并缓存
        if hasattr(wo_a, 'weight_scale_inv'):
            # fp8 权重 : 乘以 scale 然后转 bf16
            scale = wo_a.weight_scale_inv.view(n_local_groups, o_lora_rank, hidden_dim)
            cached = (wo_a.weight * scale).to(torch.bfloat16)
        else:
            # 非 fp8 权重 : 直接 reshape 并转 bf16
            cached = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to(torch.bfloat16)
        wo_a._dsv4_wo_a_bf16 = cached
    return wo_a._dsv4_wo_a_bf16
tests/kernels/attention/test_rocm_triton_attn_dsv4.py test-coverage

新增三个测试覆盖融合 kernel 正确性(与官方嵌入对比)、空张量边界、缓存功能。

# tests/kernels/attention/test_rocm_triton_attn_dsv4.py ( 片段 )@pytest.mark.parametrize('num_tokens', [1, 7, 64])
@pytest.mark.parametrize('num_heads', [1, 8])
@pytest.mark.parametrize('pos_dtype', [torch.int32, torch.int64])
@torch.inference_mode()
def test_fused_inverse_rope_gptj_matches_rotary_native(
    num_tokens: int, num_heads: int, pos_dtype: torch.dtype, default_vllm_config
) -> None:
    '''验证融合 kernel 与官方旋转嵌入的 forward_native(inverse=True) 一致。    使用了真实的 DeepseekV4ScalingRotaryEmbedding 实例,确保正确性。
    '''
    from vllm.v1.attention.ops.rocm_aiter_mla_sparse import _fused_inverse_rope_gptj
​
    device = torch.device('cuda')
    torch.manual_seed(0)
    # 构造官方旋转嵌入(小规模用于测试)
    rotary_emb = _make_dsv4_rotary(device)
    o = torch.randn(
        num_tokens, num_heads, HEAD_DIM, dtype=torch.bfloat16, device=device
    )
    positions = torch.randint(
        0, _ROTARY_CACHE_LEN, (num_tokens,), dtype=pos_dtype, device=device
    )
​
    actual = _fused_inverse_rope_gptj(
        o, positions, rotary_emb.cos_sin_cache, ROPE_HEAD_DIM
    )
    expected = _inv_rope_via_rotary_native(rotary_emb, o, positions)
​
    assert actual.dtype == torch.bfloat16
    assert actual.shape == o.shape
    # NoPE 部分透传,应为精确匹配
    assert torch.equal(actual[..., :NOPE_HEAD_DIM], expected[..., :NOPE_HEAD_DIM])
    # RoPE 部分允许 1~2 ulp 误差 ( 来自 fma 顺序 )
    torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2)

评论区精华

测试与官方旋转嵌入对比 测试

dllehr-amd 评论:单元测试是否应该使用官方 DeepseekV4ScalingRotaryEmbedding 以确保未来兼容性?

结论:作者将测试改为直接实例化官方旋转嵌入并对比 forward_native 输出。 · 已解决

移除融合逆 RoPE 的回退路径 设计

tjtanaa 指出 _can_use_fused_inv_rope 对 DSV4 始终为 True,应移除回退代码。

结论:作者删除 _can_use_fused_inv_rope 并将融合路径设为唯一路径,通过 assert 保护。 · 已解决

更新测试适配辅助函数移除 测试

tjtanaa 提醒在移除 _can_use_fused_inv_rope 后需更新单元测试。

结论:作者相应更新测试代码。 · 已解决

风险与影响

风险较低,但需注意:

  • 融合kernel假设张量布局连续、rope_dim偶数、非neox;若未来模型或配置不符,现有assert会导致运行时错误(已无回退)。但DSV4始终满足这些条件。
  • 缓存权重到模块属性,若模块被销毁或权重更新,缓存需手动清除(当前设计为惰性缓存,首次调用时填充;若训练/权重热更新可能导致状态不一致,但推理场景无影响)。
  • 仅影响ROCm上的DSV4模型,其他模型无变化。

影响范围局限在ROCm DSV4模型decode阶段:

  • 性能:并发4-64时吞吐量提升8.9%-15.6%,TPOT减少7.9%-13.3%,对延迟敏感场景收益明显。
  • 代码:只修改了两个文件,主文件减少约59行(删除回退代码),新增约126行。
  • 可维护性:融合kernel和缓存逻辑使算子更高效,但去除了回退路径,对平台兼容性要求更高。
ROCm 专用优化 无 PyTorch 回退路径 运行时 assert 保护

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论