Support online MXFP8 quantization for ungated MoE (#27939)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-06-24 16:58:48 -07:00
committed by GitHub
co-authored by Brayden Zhong
parent e4bf0043fe
commit f82addd4a8
4 changed files with 27 additions and 16 deletions
@@ -200,7 +200,7 @@ On Ascend, various layers quantization configurations are supported, see [Ascend
## GEMM Backends for FP4/FP8 Quantization
<Note>
Backend selection is supported only for **blockwise FP8** and **NVFP4** GEMM. When running FP8 or FP4 quantized models, you can select the GEMM backend via `--fp8-gemm-backend` and `--fp4-gemm-backend`.
Backend selection applies to **blockwise FP8**, **MXFP8** (dense linear), and **NVFP4** GEMM. When running offline or online FP8 or FP4 quantized models, you can select the GEMM backend via `--fp8-gemm-backend` and `--fp4-gemm-backend`.
</Note>
### `--fp8-gemm-backend` (Blockwise FP8 GEMM)
@@ -259,6 +259,8 @@ Backend selection is supported only for **blockwise FP8** and **NVFP4** GEMM. Wh
**`auto` selection order:** 1) DeepGEMM (SM90/SM100, installed); 2) FlashInfer TRTLLM (SM100, FlashInfer available); 3) CUTLASS (SM90/SM100/120); 4) AITER (AMD); 5) Triton. **Exception:** SM120 always resolves to Triton.
**MXFP8 dense linear:** `auto` uses `flashinfer_cutlass` on SM100 (else `triton`). `flashinfer_cutlass` is fastest on most shapes; `flashinfer_trtllm` is faster only at small M.
### `--fp4-gemm-backend` (NVFP4 GEMM)
<table>
@@ -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)
+6 -4
View File
@@ -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,