[AMD][DSv4] Fuse inverse-RoPE into the fp8 wo_a quant (stacked on #37423) (#37658)

Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
karverma-amd
2026-09-04 00:51:29 -07:00
committed by GitHub
co-authored by wunhuang
parent 0b57847ebf
commit dae126d510
3 changed files with 177 additions and 69 deletions
+4
View File
@@ -1402,6 +1402,10 @@ class Envs:
# cache, GEMM, and distributed
SGLANG_OPT_FP8_WO_A_GEMM = EnvBool(True)
# ROCm gfx950: fuse inverse-RoPE into the wo_a mxfp8 quant (aiter
# inverse_rope_group_quant) instead of a separate fused_rope_inplace + Triton
# quant. Off by default; requires SGLANG_OPT_FP8_WO_A_GEMM and the aiter op.
SGLANG_OPT_FP8_WO_A_FUSED_INVROPE = EnvBool(False)
# Route the decode wo_a bf16 batched matmul off rocBLAS/Tensile onto aiter's
# tuned batched_gemm_bf16 (gfx95). Off by default; see deepseek_v4.py
# _apply_wo_a_bf16_matmul.
@@ -65,12 +65,38 @@ if _is_hip and _is_gfx95_supported:
err,
)
# Fused inverse-RoPE + per-token-group mxfp8 quant in a single aiter kernel.
# Optional enhancement over the two-kernel path (fused_rope_inplace + the Triton
# quantizer above): it removes one launch and one HBM round-trip of the [T,H,Dh]
# attention output. Resolved once at import; None leaves the two-kernel path in
# place, so this is purely additive.
_inverse_rope_group_quant = None
if _batched_gemm_a8w8_mxscale is not None:
try:
from aiter.ops.inverse_rope_group_quant import (
inverse_rope_group_quant as _inverse_rope_group_quant,
)
except Exception as err: # pragma: no cover - env-dependent
logger.warning(
"aiter inverse_rope_group_quant import failed; the DSV4 wo_a fp8 "
"fused inverse-RoPE path is unavailable, using the two-kernel "
"path: %s",
err,
)
def is_wo_a_fp8_mxscale_supported() -> bool:
"""True when the ROCm fp8 ``wo_a`` path can run on this build/arch."""
return _batched_gemm_a8w8_mxscale is not None
def is_wo_a_fp8_fused_invrope_supported() -> bool:
"""True when the fused inverse-RoPE + quant ``wo_a`` path can run."""
return (
_batched_gemm_a8w8_mxscale is not None and _inverse_rope_group_quant is not None
)
@triton.jit
def _wo_a_quant_mxfp8_kernel(
x_ptr,
@@ -160,6 +186,50 @@ def apply_wo_a_fp8_mxscale(
)
def apply_wo_a_fp8_mxscale_fused_invrope(
o: torch.Tensor,
positions: torch.Tensor,
cos_cache: torch.Tensor,
sin_cache: torch.Tensor,
num_groups: int,
weight: torch.Tensor,
weight_scale: torch.Tensor,
) -> torch.Tensor:
"""fp8 ``wo_a`` with a fused inverse-RoPE + group-quant front end.
Identical result to ``apply_wo_a_fp8_mxscale`` preceded by the standalone
inverse RoPE, but ``inverse_rope_group_quant`` does the inverse rotation and
the [T,G,D] mxfp8 quant in one kernel, so the caller must NOT run the
separate ``fused_rope_inplace`` first.
Args:
o: ``[T, H, head_dim]`` bf16 attention output, still RoPE'd (the kernel
applies the inverse rotation to the trailing rope dims internally).
positions: ``[T]`` absolute positions into the cos/sin tables.
cos_cache/sin_cache: ``[max_pos, rope_dim // 2]`` bf16 tables.
num_groups: local o-groups ``G``.
weight: ``[G, R, D]`` fp8 ``wo_a`` weight.
weight_scale: ``[G, R/128, D/128]`` uint8 e8m0 (from load-time convert).
Returns bf16 ``[T, G, R]``.
"""
# Default scale layout is the unshuffled row layout on both the current
# (``scale_shuffle=False``) and older (``scale_layout="row"``) aiter
# signatures, which is exactly what ``batched_gemm_a8w8_mxscale`` consumes,
# so we omit the layout kwarg to stay compatible across aiter versions.
o_fp8, o_scale = _inverse_rope_group_quant(
o,
positions,
cos_cache,
sin_cache,
num_groups=num_groups,
quant_group_size=WO_A_MXFP8_GROUP_SIZE,
)
return _batched_gemm_a8w8_mxscale(
o_fp8, weight, o_scale, weight_scale, dtype=torch.bfloat16
)
def _is_power_of_two(scale: torch.Tensor) -> bool:
"""True when every fp32 scale is an exact power of two (zero mantissa)."""
bits = scale.detach().float().contiguous().view(torch.int32)
+103 -69
View File
@@ -392,10 +392,13 @@ _wo_a_aiter_batched_gemm_disabled = False
# the kernel availability and the weight-scale converter resolve once at import;
# ``None`` here means the platform keeps the bf16 absorb GEMM.
_wo_a_fp8_mxscale = None
_wo_a_fp8_mxscale_fused_invrope = None
_wo_a_weight_scale_to_e8m0 = None
if _is_hip:
from sglang.srt.models.deepseek_common.amd.deepseek_v4_wo_a_fp8 import (
apply_wo_a_fp8_mxscale,
apply_wo_a_fp8_mxscale_fused_invrope,
is_wo_a_fp8_fused_invrope_supported,
is_wo_a_fp8_mxscale_supported,
wo_a_weight_scale_to_e8m0,
)
@@ -403,6 +406,13 @@ if _is_hip:
if is_wo_a_fp8_mxscale_supported():
_wo_a_fp8_mxscale = apply_wo_a_fp8_mxscale
_wo_a_weight_scale_to_e8m0 = wo_a_weight_scale_to_e8m0
# Opt-in fused inverse-RoPE + quant front end (env-gated for A/B). Only
# bind it when both the flatmm and the fused aiter op are available.
if (
envs.SGLANG_OPT_FP8_WO_A_FUSED_INVROPE.get()
and is_wo_a_fp8_fused_invrope_supported()
):
_wo_a_fp8_mxscale_fused_invrope = apply_wo_a_fp8_mxscale_fused_invrope
def _apply_wo_a_bf16_matmul(
@@ -1757,83 +1767,107 @@ class MQALayer(MqaAttentionBase):
save_kv_cache=save_kv_cache,
)
o = o[:, tp_slice, :]
if _is_npu:
cos4, sin4 = self._get_npu_rope_position_cache(
positions, o.dtype, inverse=True
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
if (
_FP8_WO_A_GEMM
and _wo_a_fp8_mxscale_fused_invrope is not None
and not _is_npu
):
# ROCm gfx950 fused path: inverse-RoPE + per-token-group mxfp8 quant
# in one aiter kernel on the pre-view [T,H,Dh] output, then the a8w8
# mxscale absorb GEMM. Replaces the standalone inverse RoPE, the
# [T,G,D] view, and the quant inside the two-kernel fp8 path below.
G = self.n_local_groups
cos_c, sin_c = _freqs_cis_to_cos_sin(self.freqs_cis, o.dtype, o.device)
o = _wo_a_fp8_mxscale_fused_invrope(
o,
None,
cos4,
sin4,
qk_nope_dim=self.qk_nope_head_dim,
)
else:
fused_rope_inplace(
o[..., -self.qk_rope_head_dim :],
None,
self.freqs_cis,
positions=positions,
inverse=True,
)
o = o.view(o.shape[0], self.n_local_groups, -1)
if _FP8_WO_A_GEMM and _wo_a_fp8_mxscale is not None:
# ROCm gfx950: same fp8 absorb GEMM as the DeepGEMM path below, but
# through aiter's e8m0 block-scale batched GEMM. The activation is
# quantized per token-group inside the helper.
T, G, D = o.shape
o = _wo_a_fp8_mxscale(
o,
self.wo_a.weight.view(G, self.o_lora_rank, D),
positions,
cos_c,
sin_c,
G,
self.wo_a.weight.view(G, self.o_lora_rank, -1),
self.wo_a.weight_scale_inv.data,
)
elif _FP8_WO_A_GEMM:
import deep_gemm
from sglang.srt.layers import deep_gemm_wrapper
T, G, D = o.shape
R = self.o_lora_rank
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# sm100 (Blackwell): ue8m0 scales via the dedicated JIT kernel.
o_fp8, o_s = sglang_per_token_group_quant_fp8_dsv4_wo_a(o)
recipe = (1, 1, 128)
else:
# sm90 (Hopper): fp32 scales.
o_fp8, o_s = sglang_per_token_group_quant_fp8(
o.reshape(T * G, D).contiguous(),
group_size=128,
scale_ue8m0=False,
)
o_fp8 = o_fp8.view(T, G, D)
o_s = o_s.view(T, G, -1)
recipe = (1, 128, 128)
output = torch.empty(T, G, R, device=o.device, dtype=torch.bfloat16)
deep_gemm.fp8_einsum(
"bhr,hdr->bhd",
(o_fp8, o_s),
(self.wo_a.weight.view(G, R, D), self.wo_a.weight_scale_inv.data),
output,
recipe=recipe,
)
o = output
else:
wo_a_weight = getattr(self.wo_a, "weight", None)
if wo_a_weight is not None:
wo_a = wo_a_weight.view(self.n_local_groups, self.o_lora_rank, -1)
o = _apply_wo_a_bf16_matmul(
o, wo_a, is_decode=forward_batch.forward_mode.is_decode()
if _is_npu:
cos4, sin4 = self._get_npu_rope_position_cache(
positions, o.dtype, inverse=True
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
o,
None,
cos4,
sin4,
qk_nope_dim=self.qk_nope_head_dim,
)
else:
o = _apply_gguf_grouped_wo_a(
o,
self.wo_a.qweight,
self.wo_a.qweight_type.weight_type,
self.o_lora_rank,
fused_rope_inplace(
o[..., -self.qk_rope_head_dim :],
None,
self.freqs_cis,
positions=positions,
inverse=True,
)
o = o.view(o.shape[0], self.n_local_groups, -1)
if _FP8_WO_A_GEMM and _wo_a_fp8_mxscale is not None:
# ROCm gfx950: same fp8 absorb GEMM as the DeepGEMM path below,
# but through aiter's e8m0 block-scale batched GEMM. The
# activation is quantized per token-group inside the helper.
T, G, D = o.shape
o = _wo_a_fp8_mxscale(
o,
self.wo_a.weight.view(G, self.o_lora_rank, D),
self.wo_a.weight_scale_inv.data,
)
elif _FP8_WO_A_GEMM:
import deep_gemm
from sglang.srt.layers import deep_gemm_wrapper
T, G, D = o.shape
R = self.o_lora_rank
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# sm100 (Blackwell): ue8m0 scales via the dedicated JIT kernel.
o_fp8, o_s = sglang_per_token_group_quant_fp8_dsv4_wo_a(o)
recipe = (1, 1, 128)
else:
# sm90 (Hopper): fp32 scales.
o_fp8, o_s = sglang_per_token_group_quant_fp8(
o.reshape(T * G, D).contiguous(),
group_size=128,
scale_ue8m0=False,
)
o_fp8 = o_fp8.view(T, G, D)
o_s = o_s.view(T, G, -1)
recipe = (1, 128, 128)
output = torch.empty(T, G, R, device=o.device, dtype=torch.bfloat16)
deep_gemm.fp8_einsum(
"bhr,hdr->bhd",
(o_fp8, o_s),
(
self.wo_a.weight.view(G, R, D),
self.wo_a.weight_scale_inv.data,
),
output,
recipe=recipe,
)
o = output
else:
wo_a_weight = getattr(self.wo_a, "weight", None)
if wo_a_weight is not None:
wo_a = wo_a_weight.view(self.n_local_groups, self.o_lora_rank, -1)
o = _apply_wo_a_bf16_matmul(
o, wo_a, is_decode=forward_batch.forward_mode.is_decode()
)
else:
o = _apply_gguf_grouped_wo_a(
o,
self.wo_a.qweight,
self.wo_a.qweight_type.weight_type,
self.o_lora_rank,
)
o, _ = self.wo_b(o.flatten(1))
if self.attn_tp_size > 1 and self.attn_tp_size < get_parallel().tp_size:
o = attn_tp_all_reduce(o)