[ROCm] dsv4: remove the redundant fp8 scale transpose-copy on decode (#27289)
Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com> Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
co-authored by
Zhiyao Jiang
Thomas Wang
parent
fca4ef9d69
commit
ea1d190ed0
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user