[AMD] Use bpreshuffle FP8 blockscale GEMM to replace ABScale GEMM (#23319)

Co-authored-by: HaiShaw <hixiao@gmail.com>
This commit is contained in:
Xinyi Song
2026-04-23 01:51:30 -07:00
committed by GitHub
co-authored by HaiShaw
parent fd88a1c562
commit cd459af4e2
2 changed files with 39 additions and 13 deletions
@@ -50,6 +50,7 @@ from sglang.srt.layers.quantization.fp8_kernel import (
scaled_fp8_quant,
)
from sglang.srt.layers.quantization.fp8_utils import (
_use_aiter_gfx95,
apply_fp8_linear,
can_auto_enable_marlin_fp8,
cutlass_fp8_supported,
@@ -114,6 +115,12 @@ if _use_aiter or _use_hip_int4:
from aiter.fused_moe import fused_moe
from aiter.ops.shuffle import shuffle_weight
if _use_aiter:
from sglang.srt.layers.quantization.fp8_utils import (
aiter_w8a8_block_fp8_linear,
use_aiter_triton_gemm_w8a8_tuned_gfx950,
)
ACTIVATION_SCHEMES = ["static", "dynamic"]
@@ -502,6 +509,18 @@ class Fp8LinearMethod(LinearMethodBase):
layer.weight.data = weight.data
layer.weight_scale_inv.data = weight_scale.data
if (
_use_aiter_gfx95
and self.w8a8_block_fp8_linear is aiter_w8a8_block_fp8_linear
):
n, k = layer.weight.shape
if not use_aiter_triton_gemm_w8a8_tuned_gfx950(n, k):
# TODO(1am9trash), to deal with case that this branch chance
# drops as use_aiter_triton_gemm_w8a8_tuned_gfx950() expands
t = shuffle_weight(layer.weight, (16, 16))
layer.weight.copy_(t)
del t
def _process_mxfp8_linear_weight_scale(self, layer: Module) -> None:
if not self.use_mxfp8:
return
@@ -80,10 +80,11 @@ def use_aiter_triton_gemm_w8a8_tuned_gfx950(n: int, k: int) -> bool:
if _use_aiter:
import aiter
# from aiter import gemm_a8w8_blockscale, gemm_a8w8_bpreshuffle, get_hip_quant
from aiter import gemm_a8w8_blockscale as gemm_a8w8_blockscale
from aiter import gemm_a8w8_bpreshuffle, get_hip_quant
from aiter import (
gemm_a8w8_blockscale_bpreshuffle,
gemm_a8w8_bpreshuffle,
get_hip_quant,
)
from aiter.ops.triton.gemm_a8w8_blockscale import (
gemm_a8w8_blockscale as triton_gemm_a8w8_blockscale,
)
@@ -756,14 +757,6 @@ def aiter_w8a8_block_fp8_linear(
input_2d = input.view(-1, input.shape[-1])
output_shape = [*input.shape[:-1], weight.shape[0]]
# if input_scale not None, input is quanted
if input_scale is not None:
q_input = input_2d
x_scale = input_scale
else:
q_input, x_scale = aiter_per1x128_quant(input_2d, quant_dtype=aiter.dtypes.fp8)
n, k = weight.shape
if _use_aiter_gfx95:
@@ -771,10 +764,24 @@ def aiter_w8a8_block_fp8_linear(
else:
use_triton = True
# if input_scale not None, input is quanted
if input_scale is not None:
q_input = input_2d
x_scale = input_scale
if not use_triton:
x_scale = x_scale.transpose(-1, -2).contiguous().view(*x_scale.shape)
else:
q_input, x_scale = aiter_per1x128_quant(
input_2d,
quant_dtype=aiter.dtypes.fp8,
transpose_scale=not use_triton,
)
if use_triton:
gemm_a8w8_blockscale_op = triton_gemm_a8w8_blockscale
else:
gemm_a8w8_blockscale_op = gemm_a8w8_blockscale
# TODO(1am9trash), to deal with chance of this branch changes
gemm_a8w8_blockscale_op = gemm_a8w8_blockscale_bpreshuffle
output = gemm_a8w8_blockscale_op(
q_input,