执行摘要
- 一句话:融合逆RoPE与wo_a缓存,ROCm DSV4 decode性能提升9-16%
- 推荐动作:值得精读,该PR展示了如何通过Triton kernel融合和权重缓存实现显著性能优化,设计思路清晰,测试充分。尤其是将多步小kernel合并为单launch、以及惰性缓存模式的实践,对其他算子优化有参考价值。
功能与动机
从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性能瓶颈。
实现拆解
- 引入Triton融合kernel:在
rocm_aiter_mla_sparse.py中定义_inverse_rope_gptj_kernel(@triton.jit),一次launch完成逆RoPE,直接输出bf16。该kernel假设非neox布局、rope_dim为偶数、连续张量。
- 替换逆RoPE调用:将
rocm_inv_rope_einsum中的_apply_inv_rope_ref(包含PyTorch回退)替换为_fused_inverse_rope_gptj调用,并根据review建议移除了_can_use_fused_inv_rope函数,使融合路径成为唯一路径(通过assert确保条件满足)。
- 缓存反量化wo_a权重:在
_get_cached_wo_a_bf16中,首次调用时将fp8权重反量化并存储为模块属性wo_a._dsv4_wo_a_bf16,后续步骤直接返回缓存,避免每步重复反量化。
- 添加配套测试:在
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(模块 注意力层;类别 source;类型 core-logic;符号 _apply_gptj_inv_rope_ref, _inverse_rope_gptj_kernel, _fused_inverse_rope_gptj, _apply_inv_rope_ref): 核心优化实现文件,包含Triton融合kernel、逆RoPE替换、wo_a缓存逻辑。
tests/kernels/attention/test_rocm_triton_attn_dsv4.py(模块 测试;类别 test;类型 test-coverage;符号 _make_dsv4_rotary, _inv_rope_via_rotary_native, _FakeWoA, init): 新增三个测试覆盖融合kernel正确性(与官方嵌入对比)、空张量边界、缓存功能。
关键符号:_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
核心优化实现文件,包含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
新增三个测试覆盖融合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)
评论区精华
核心review讨论集中于:
风险与影响
- 风险:风险较低,但需注意:
- 融合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保护
关联脉络
- PR #39457 [V1][Metrics] Add MLA attention metrics for DeepSeek MFU estimation: 同属DeepSeek模型在V1引擎上的优化,本PR关注o-projection性能,39457关注指标测量,共同支撑DeepSeek在ROCm上的部署和分析。
参与讨论