diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 1ca9deab7..db8c0210b 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -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 diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 9c578ee41..8f3c84bf8 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -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,