From ea1d190ed02611787c8fb2ea2ba76a394dca9111 Mon Sep 17 00:00:00 2001 From: Xinyu Jiang Date: Mon, 8 Jun 2026 14:49:34 -0400 Subject: [PATCH] [ROCm] dsv4: remove the redundant fp8 scale transpose-copy on decode (#27289) Co-authored-by: Zhiyao Jiang Co-authored-by: Thomas Wang --- python/sglang/srt/layers/communicator.py | 3 +++ python/sglang/srt/layers/quantization/fp8_utils.py | 6 ++++-- .../attention_forward_methods/forward_mha.py | 4 ++++ .../attention_forward_methods/forward_mla.py | 3 +++ python/sglang/srt/models/deepseek_common/utils.py | 2 ++ python/sglang/srt/models/deepseek_v2.py | 3 ++- python/sglang/srt/models/deepseek_v4.py | 2 ++ 7 files changed, 20 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 7a482f999..b428adafc 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -65,6 +65,7 @@ from sglang.srt.layers.moe import ( should_use_dp_reduce_scatterv, should_use_flashinfer_cutlass_moe_fp4_allgather, ) +from sglang.srt.layers.quantization.fp8_utils import _use_aiter_bpreshuffle_gfx95 from sglang.srt.layers.utils.cp_utils import ( is_mla_prefill_cp_enabled, mla_use_prefill_cp, @@ -572,6 +573,7 @@ class LayerCommunicator: dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=_dsa_needs_bf16, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) if _dsa_needs_bf16: hidden_states = ( @@ -617,6 +619,7 @@ class LayerCommunicator: dtype_quant=torch.float8_e4m3fn, res1=residual, output_unquantized_inp1=_dsa_needs_bf16, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) ) if _dsa_needs_bf16: diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 51dc933d1..6c6827c7d 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -786,8 +786,10 @@ def aiter_w8a8_block_fp8_linear( if input_scale is not None: q_input = input_2d x_scale = input_scale - if _use_aiter_bpreshuffle_gfx95 and not use_triton: - x_scale = x_scale.transpose(-1, -2).contiguous().view(*x_scale.shape) + # On ROCm >= 7.2, scale is in bpreshuffle's transposed layout. + # Triton needs a row-major view, so adjust strides only. No copy. + if use_triton and _use_aiter_bpreshuffle_gfx95: + x_scale = torch.as_strided(x_scale, x_scale.shape, (1, x_scale.shape[0])) else: q_input, x_scale = aiter_per1x128_quant( input_2d, diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index dbcb3ee0f..729be4904 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -19,6 +19,7 @@ from sglang.srt.models.deepseek_common.utils import ( _is_hip, _is_musa, _is_npu, + _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) from sglang.srt.server_args import get_global_server_args @@ -152,6 +153,7 @@ class DeepseekMHAForwardMixin: dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=True, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) q = self.q_b_proj(q_quanted)[0].view( -1, self.num_local_heads, self.qk_head_dim @@ -193,6 +195,7 @@ class DeepseekMHAForwardMixin: dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=False, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) else: @@ -222,6 +225,7 @@ class DeepseekMHAForwardMixin: dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=True, # return unqaunt kv_a + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) else: diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 54affa654..d9fcbf807 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -38,6 +38,7 @@ from sglang.srt.models.deepseek_common.utils import ( _is_hip, _is_musa, _use_aiter, + _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) from sglang.srt.server_args import get_global_server_args @@ -197,6 +198,7 @@ class DeepseekMLAForwardMixin: dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=True, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) q = q_quanted else: @@ -211,6 +213,7 @@ class DeepseekMLAForwardMixin: dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=False, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) elif _use_aiter: diff --git a/python/sglang/srt/models/deepseek_common/utils.py b/python/sglang/srt/models/deepseek_common/utils.py index 418e65959..a23c13978 100644 --- a/python/sglang/srt/models/deepseek_common/utils.py +++ b/python/sglang/srt/models/deepseek_common/utils.py @@ -25,6 +25,7 @@ from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, get_device_sm, + get_hip_version, is_cpu, is_cuda, is_gfx95_supported, @@ -47,6 +48,7 @@ _is_xpu = is_xpu() _device_sm = get_device_sm() _is_gfx95_supported = is_gfx95_supported() _use_aiter_gfx95 = _use_aiter and _is_gfx95_supported +_use_aiter_bpreshuffle_gfx95 = _use_aiter_gfx95 and get_hip_version() >= (7, 2, 0) _is_cublas_ge_129 = is_nvidia_cublas_version_ge_12_9() diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index c5fc1113f..c34681ced 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -162,6 +162,7 @@ from sglang.srt.models.deepseek_common.utils import ( _is_npu, _is_xpu, _use_aiter, + _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) from sglang.srt.server_args import get_global_server_args @@ -376,7 +377,7 @@ class DeepseekV2MLP(nn.Module): swiglu_limit=self.swiglu_limit, activation="silu", dtype_quant=dtypes.fp8, - transpose_scale=False, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) x = (x_fp8, x_scale) else: diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 7786c9edb..0594829bb 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -97,6 +97,7 @@ from sglang.srt.models.dbrx import ReplicatedLinear from sglang.srt.models.deepseek_common.amd.deepseek_v4_fused_mhc import ( try_fused_hc_post_pre, ) +from sglang.srt.models.deepseek_common.utils import _use_aiter_bpreshuffle_gfx95 from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu if not _is_hip: @@ -151,6 +152,7 @@ def _fused_rmsnorm_fp8_quant(hidden_states, weight, eps): dtype_quant=torch.float8_e4m3fn, res1=None, output_unquantized_inp1=True, + transpose_scale=_use_aiter_bpreshuffle_gfx95, ) return x_quant, x_bf16