[NPU] Fuse MXFP4 W4A8 MoE gmm1 + swiglu + requant into one kernel (#39881)
This commit is contained in:
@@ -282,6 +282,16 @@ class NPUW4A8MXFP4MoEMethod(_NPUMoEMethodBase):
|
|||||||
def __init__(self, dynamic_quant_kwargs=_DEFAULT_DYNAMIC_QUANT):
|
def __init__(self, dynamic_quant_kwargs=_DEFAULT_DYNAMIC_QUANT):
|
||||||
super().__init__(quant_config=None)
|
super().__init__(quant_config=None)
|
||||||
self.dynamic_quant_kwargs = dynamic_quant_kwargs
|
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(
|
def process_weights_after_loading(
|
||||||
self, layer: torch.nn.Module, weight_prefix: str
|
self, layer: torch.nn.Module, weight_prefix: str
|
||||||
@@ -340,6 +350,57 @@ class NPUW4A8MXFP4MoEMethod(_NPUMoEMethodBase):
|
|||||||
dynamic_quant_kwargs=dynamic_quant_kwargs,
|
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
|
# NPUW4A4MXFP4MoEMethod
|
||||||
|
|||||||
@@ -26,6 +26,15 @@ from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
|
|||||||
NPUW4A8MXFP4MoEMethod,
|
NPUW4A8MXFP4MoEMethod,
|
||||||
NPUW8A8Int8MoEMethod,
|
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 (
|
from sglang.srt.layers.moe.moe_runner.base import (
|
||||||
MoeQuantInfo,
|
MoeQuantInfo,
|
||||||
MoeRunnerConfig,
|
MoeRunnerConfig,
|
||||||
@@ -93,13 +102,15 @@ class AscendRunnerCore(MoeRunnerCore):
|
|||||||
|
|
||||||
kernel = config.layer.w2_kernel
|
kernel = config.layer.w2_kernel
|
||||||
|
|
||||||
if isinstance(kernel, NPUMXFP8MoEMethod):
|
if _uses_fused_gmm1(kernel):
|
||||||
# MXFP8 fuses gate/up + swiglu + requant into gmm1, so there is no
|
# Fused methods (MXFP8; MXFP4 W4A8 via use_fused_gmm1) fold
|
||||||
# separate activation step — run() skips it. Left None on purpose so
|
# gate/up + swiglu + requant into gmm1, so there is no separate
|
||||||
# that reaching for it fails loudly instead of silently applying an
|
# activation step — run() skips it. Left None on purpose so that
|
||||||
# unfused swiglu to already-requantised activations. This holds for
|
# reaching for it fails loudly instead of silently applying an
|
||||||
# both dispatchers: ascend_tp gets its activation quant fused into
|
# unfused swiglu to already-requantised activations. This holds
|
||||||
# routing, DeepEP dispatches bf16 and gmm1 quantises it itself.
|
# for both dispatchers: ascend_tp gets its activation quant fused
|
||||||
|
# into routing, DeepEP dispatches bf16 and gmm1 quantises it
|
||||||
|
# itself.
|
||||||
self.activation = None
|
self.activation = None
|
||||||
elif (
|
elif (
|
||||||
isinstance(kernel, NPUW4A8MXFP4MoEMethod)
|
isinstance(kernel, NPUW4A8MXFP4MoEMethod)
|
||||||
@@ -190,10 +201,10 @@ class AscendRunnerCore(MoeRunnerCore):
|
|||||||
|
|
||||||
w13_kernel = self.config.layer.w13_kernel
|
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 ---
|
# --- w13 projection + activation, fused into one kernel ---
|
||||||
# MXFP8 gmm1 returns activations already requantised for gmm2, so
|
# The fused gmm1 returns activations already requantised for gmm2,
|
||||||
# there is no separate activation step to run.
|
# so there is no separate activation step to run.
|
||||||
hidden_states, pertoken_scale = w13_kernel.apply_fused_gmm1_swiglu(
|
hidden_states, pertoken_scale = w13_kernel.apply_fused_gmm1_swiglu(
|
||||||
quant_info,
|
quant_info,
|
||||||
x,
|
x,
|
||||||
|
|||||||
Reference in New Issue
Block a user