Support online MXFP8 quantization for ungated MoE (#27939)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
e4bf0043fe
commit
f82addd4a8
@@ -322,6 +322,11 @@ def align_mxfp8_moe_weights_for_flashinfer_trtllm(layer: Module) -> None:
|
||||
assert w13_scale.dtype == torch.uint8
|
||||
assert w2_scale.dtype == torch.uint8
|
||||
|
||||
if not is_gated:
|
||||
intermediate = w2_weight.shape[2]
|
||||
w13_weight = w13_weight[:, :intermediate, :].contiguous()
|
||||
w13_scale = w13_scale[:, :intermediate, :].contiguous()
|
||||
|
||||
# Pad for kernel alignment (non-gated needs 128, gated needs 16)
|
||||
min_alignment = 16 if is_gated else 128
|
||||
w13_weight, w13_scale, w2_weight, w2_scale, _ = _align_mxfp8_moe_weights(
|
||||
@@ -675,7 +680,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
|
||||
assert quant_info.weight_block_k == 32
|
||||
from flashinfer import mxfp8_quantize
|
||||
|
||||
a_q, a_sf = mxfp8_quantize(hidden_states, False)
|
||||
a_q, a_sf = mxfp8_quantize(hidden_states, False, backend="cute-dsl")
|
||||
# FlashInfer TRT-LLM MxFP8 expects token-major activation scales:
|
||||
# [num_tokens, hidden_size // 32] (no transpose).
|
||||
a_sf_t = a_sf.view(torch.uint8).reshape(hidden_states.shape[0], -1)
|
||||
|
||||
@@ -587,7 +587,8 @@ class Fp8LinearMethod(LinearMethodBase):
|
||||
if not self.use_mxfp8:
|
||||
return
|
||||
|
||||
if get_fp8_gemm_runner_backend().is_flashinfer_trtllm():
|
||||
backend = get_fp8_gemm_runner_backend()
|
||||
if backend.is_flashinfer_trtllm():
|
||||
from flashinfer import shuffle_matrix_a, shuffle_matrix_sf_a
|
||||
|
||||
weight = layer.weight.data
|
||||
@@ -628,7 +629,7 @@ class Fp8LinearMethod(LinearMethodBase):
|
||||
.reshape_as(scale_u8)
|
||||
.contiguous(),
|
||||
)
|
||||
elif get_fp8_gemm_runner_backend().is_flashinfer_cutlass():
|
||||
elif backend.is_flashinfer_cutlass():
|
||||
from flashinfer import block_scale_interleave
|
||||
|
||||
scale_u8 = layer.weight_scale_inv.data
|
||||
@@ -791,9 +792,10 @@ class Fp8LinearMethod(LinearMethodBase):
|
||||
)
|
||||
|
||||
if self.use_mxfp8:
|
||||
if get_fp8_gemm_runner_backend().is_flashinfer_cutlass():
|
||||
backend = get_fp8_gemm_runner_backend()
|
||||
if backend.is_flashinfer_cutlass():
|
||||
weight_scale = layer.weight_scale_inv_swizzled
|
||||
elif get_fp8_gemm_runner_backend().is_flashinfer_trtllm():
|
||||
elif backend.is_flashinfer_trtllm():
|
||||
weight_scale = layer.weight_scale_inv_shuffled
|
||||
else:
|
||||
weight_scale = layer.weight_scale_inv
|
||||
|
||||
@@ -410,15 +410,8 @@ def dispatch_w8a8_block_fp8_linear() -> Callable:
|
||||
|
||||
|
||||
def dispatch_w8a8_mxfp8_linear() -> Callable:
|
||||
"""Dispatch MXFP8 linear kernel by --fp8-gemm-backend.
|
||||
|
||||
For MXFP8, Triton remains the default path. We only route to FlashInfer
|
||||
when backend is explicitly set to flashinfer_cutlass or flashinfer_trtllm.
|
||||
"""
|
||||
backend = get_fp8_gemm_runner_backend()
|
||||
if backend.is_flashinfer_trtllm():
|
||||
return flashinfer_mxfp8_blockscaled_linear
|
||||
elif backend.is_flashinfer_cutlass():
|
||||
if backend.is_flashinfer_cutlass() or backend.is_flashinfer_trtllm():
|
||||
return flashinfer_mxfp8_blockscaled_linear
|
||||
return triton_mxfp8_blockscaled_linear
|
||||
|
||||
@@ -516,7 +509,17 @@ def initialize_fp8_gemm_config(server_args: ServerArgs) -> None:
|
||||
# TODO(brayden): Verify if CUTLASS can be set by default once SwapAB is supported
|
||||
backend = "triton"
|
||||
|
||||
FP8_GEMM_RUNNER_BACKEND = Fp8GemmRunnerBackend(backend)
|
||||
backend = Fp8GemmRunnerBackend(backend)
|
||||
|
||||
if (
|
||||
backend.is_auto()
|
||||
and server_args.quantization == "mxfp8"
|
||||
and _is_sm100_supported
|
||||
and is_flashinfer_available()
|
||||
):
|
||||
backend = Fp8GemmRunnerBackend.FLASHINFER_CUTLASS
|
||||
|
||||
FP8_GEMM_RUNNER_BACKEND = backend
|
||||
|
||||
|
||||
def get_fp8_gemm_runner_backend() -> Fp8GemmRunnerBackend:
|
||||
@@ -1121,7 +1124,6 @@ def flashinfer_mxfp8_blockscaled_linear(
|
||||
weight_t = weight.contiguous().t()
|
||||
|
||||
if get_fp8_gemm_runner_backend().is_flashinfer_trtllm():
|
||||
|
||||
weight_scale_t = weight_scale.contiguous().view(-1)
|
||||
output = flashinfer_mm_mxfp8(
|
||||
q_input,
|
||||
|
||||
Reference in New Issue
Block a user