[NPU] Fuse MXFP4 W4A8 MoE gmm1 + swiglu + requant into one kernel (#39881)

This commit is contained in:
iridiumine
2026-09-18 16:50:49 +08:00
committed by GitHub
parent 2dee23a876
commit d6090f92bf
2 changed files with 82 additions and 10 deletions
@@ -282,6 +282,16 @@ class NPUW4A8MXFP4MoEMethod(_NPUMoEMethodBase):
def __init__(self, dynamic_quant_kwargs=_DEFAULT_DYNAMIC_QUANT):
super().__init__(quant_config=None)
self.dynamic_quant_kwargs = dynamic_quant_kwargs
# Fused gmm1 (matmul + swiglu + requant in one aclnn kernel). The v2 op
# accepts FP4 weights via weight_scale/weight_dtype — verified on A5
# (llm/probe_mxfp4_gmm_swiglu_quant.py: same numerics as the unfused
# chain modulo the output fp8 requant, which the unfused chain also
# applies before gmm2).
self.use_fused_gmm1 = True
self.fused_matmul = GroupedMatmulSwigluQuant()
self.hidden_states_quantizer = HiddenStatesDynamicQuant(
quant_dtype=torch.float8_e4m3fn
)
def process_weights_after_loading(
self, layer: torch.nn.Module, weight_prefix: str
@@ -340,6 +350,57 @@ class NPUW4A8MXFP4MoEMethod(_NPUMoEMethodBase):
dynamic_quant_kwargs=dynamic_quant_kwargs,
)
def apply_fused_gmm1_swiglu(
self,
quant_info: "AscendQuantInfo",
hidden_states: torch.Tensor,
expert_tokens: torch.Tensor,
pertoken_scale: Optional[torch.Tensor],
group_list_type,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""gmm1 + swiglu + requant in one kernel (fused MXFP4 path).
Mirrors NPUMXFP8MoEMethod.apply_fused_gmm1_swiglu but passes the FP4
weight dtype and the e8m0 weight_scale. Returns (e4m3 activations,
e8m0 block scale) ready for the w2 gmm.
The dispatcher hands over BF16 (see process_weights_after_loading), so
pertoken_scale is normally None and the activation quant happens here.
"""
fp4_dtype = _get_float4_e2m1fn_x2_dtype()
if fp4_dtype is None:
raise RuntimeError("NPU W4A8 MXFP MoE requires float4 support.")
e8m0_dtype = _require_e8m0_dtype()
if pertoken_scale is None:
hidden_states, pertoken_scale = self.hidden_states_quantizer(hidden_states)
else:
# flat [T, K//32] -> pair form [T, K//64, 2]; identity if already
# pair-split.
pertoken_scale = pertoken_scale.reshape(
hidden_states.shape[0], hidden_states.shape[1] // 64, 2
)
return self.fused_matmul.forward(
quant_info,
"w13",
hidden_states,
expert_tokens.to(torch.int64),
group_list_type=group_list_type,
transposed=True,
weight_scale=[quant_info.w13_weight_scale],
x_scale=pertoken_scale,
dequant_mode=2,
quant_mode=2,
dequant_dtype=torch.float32,
quant_dtype=torch.float8_e4m3fn,
# e4m3 is implicit for x; FP4 must be passed for the weight.
x_dtype=None,
weight_dtype=fp4_dtype,
weight_scale_dtype=e8m0_dtype,
x_scale_dtype=e8m0_dtype,
)
# ---------------------------------------------------------------------------
# NPUW4A4MXFP4MoEMethod
@@ -26,6 +26,15 @@ from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
NPUW4A8MXFP4MoEMethod,
NPUW8A8Int8MoEMethod,
)
def _uses_fused_gmm1(kernel) -> bool:
"""Whether gmm1 runs matmul+swiglu+requant fused (no separate activation)."""
if isinstance(kernel, NPUMXFP8MoEMethod):
return True
return isinstance(kernel, NPUW4A8MXFP4MoEMethod) and kernel.use_fused_gmm1
from sglang.srt.layers.moe.moe_runner.base import (
MoeQuantInfo,
MoeRunnerConfig,
@@ -93,13 +102,15 @@ class AscendRunnerCore(MoeRunnerCore):
kernel = config.layer.w2_kernel
if isinstance(kernel, NPUMXFP8MoEMethod):
# MXFP8 fuses gate/up + swiglu + requant into gmm1, so there is no
# separate activation step — run() skips it. Left None on purpose so
# that reaching for it fails loudly instead of silently applying an
# unfused swiglu to already-requantised activations. This holds for
# both dispatchers: ascend_tp gets its activation quant fused into
# routing, DeepEP dispatches bf16 and gmm1 quantises it itself.
if _uses_fused_gmm1(kernel):
# Fused methods (MXFP8; MXFP4 W4A8 via use_fused_gmm1) fold
# gate/up + swiglu + requant into gmm1, so there is no separate
# activation step — run() skips it. Left None on purpose so that
# reaching for it fails loudly instead of silently applying an
# unfused swiglu to already-requantised activations. This holds
# for both dispatchers: ascend_tp gets its activation quant fused
# into routing, DeepEP dispatches bf16 and gmm1 quantises it
# itself.
self.activation = None
elif (
isinstance(kernel, NPUW4A8MXFP4MoEMethod)
@@ -190,10 +201,10 @@ class AscendRunnerCore(MoeRunnerCore):
w13_kernel = self.config.layer.w13_kernel
if isinstance(w13_kernel, NPUMXFP8MoEMethod):
if _uses_fused_gmm1(w13_kernel):
# --- w13 projection + activation, fused into one kernel ---
# MXFP8 gmm1 returns activations already requantised for gmm2, so
# there is no separate activation step to run.
# The fused gmm1 returns activations already requantised for gmm2,
# so there is no separate activation step to run.
hidden_states, pertoken_scale = w13_kernel.apply_fused_gmm1_swiglu(
quant_info,
x,