Skip to content

Commit bb2666d

Browse files
authored
[Feature] Add rmsnorm dynamic mx quant fusion pass (vllm-project#10730)
### What this PR does / why we need it? This PR introduces two new patterns to support the fusion of AddRMSNorm and DynamicMxQuant operators, and two new patterns to support the fusion of RMSNorm and DynamicMxQuant operators. After replacing the fusion operators, the avg execution time drops from 7.541 µs to 3.525 µs for the first combination, and from 8.098 µs to 4.625 µs for the second, and model accuracy remains unaffected. Please note that the fused operators introduced in this PR are only supported on A5. The validation was performed with CANN 9.1.T560.B030 and FrameworkPTAdapter 26.1.0.B060. ### Does this PR introduce _any_ user-facing change? N/A ### How was this patch tested? - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@967c5c3 --------- Signed-off-by: Sunwish <isunwish@foxmail.com>
1 parent c3ac524 commit bb2666d

2 files changed

Lines changed: 222 additions & 0 deletions

File tree

vllm_ascend/compilation/passes/norm_quant_fusion_pass.py

Lines changed: 200 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,10 @@
2323
from vllm.logger import logger
2424

2525
from vllm_ascend.compilation.passes.base_pattern import BasePattern
26+
from vllm_ascend.device.mxfp_compat import (
27+
is_add_rms_norm_dynamic_mx_quant_fusion_available,
28+
is_rms_norm_dynamic_mx_quant_fusion_available,
29+
)
2630
from vllm_ascend.utils import enable_custom_op
2731

2832

@@ -474,6 +478,183 @@ def replacement(
474478
return replacement
475479

476480

481+
class AddRMSNormDynamicMXQuantPattern(BasePattern):
482+
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
483+
super().__init__(vllm_config, eps)
484+
485+
def get_inputs(self):
486+
"""
487+
Generate example inputs for the AddRMSNormDynamicMXQuant fusion pattern.
488+
"""
489+
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
490+
residual = torch.randn(2, 64, device="npu", dtype=self.dtype)
491+
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
492+
return [rms_norm_input, residual, rms_norm_weight]
493+
494+
def get_pattern(self):
495+
def pattern(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
496+
"""
497+
Pattern for AddRMSNormDynamicMXQuant fusion.
498+
"""
499+
output = torch.ops.npu.npu_add_rms_norm(rms_norm_input, residual, rms_norm_weight, self.eps)
500+
out0 = output[0]
501+
out1 = output[2]
502+
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
503+
return quantized_output[0], quantized_output[1], out1
504+
505+
return pattern
506+
507+
def get_replacement(self):
508+
def replacement(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
509+
"""
510+
Replacement for the AddRMSNormDynamicMXQuant fusion.
511+
"""
512+
output = torch.ops.npu.npu_add_rms_norm_dynamic_mx_quant(
513+
rms_norm_input,
514+
residual,
515+
rms_norm_weight,
516+
epsilon=self.eps,
517+
dst_type=torch.float8_e4m3fn,
518+
)
519+
return (
520+
output[0],
521+
output[2],
522+
output[1],
523+
)
524+
525+
return replacement
526+
527+
528+
class AddRMSNormDynamicMXQuantSPPattern(BasePattern):
529+
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
530+
super().__init__(vllm_config, eps)
531+
532+
def get_inputs(self):
533+
"""
534+
Generate example inputs for the AddRMSNormDynamicMXQuant fusion pattern.
535+
"""
536+
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
537+
residual = torch.randn(2, 64, device="npu", dtype=self.dtype)
538+
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
539+
return [rms_norm_input, residual, rms_norm_weight]
540+
541+
def get_pattern(self):
542+
def pattern(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
543+
"""
544+
Pattern for AddRMSNormDynamicMXQuant fusion.
545+
"""
546+
output = torch.ops.npu.npu_add_rms_norm(rms_norm_input, residual, rms_norm_weight, self.eps)
547+
out0 = output[0]
548+
out1 = output[2]
549+
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
550+
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
551+
return quantized_output[0], quantized_output[1], out1
552+
553+
return pattern
554+
555+
def get_replacement(self):
556+
def replacement(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
557+
"""
558+
Replacement for the AddRMSNormDynamicMXQuant fusion.
559+
"""
560+
output = torch.ops.npu.npu_add_rms_norm_dynamic_mx_quant(
561+
rms_norm_input,
562+
residual,
563+
rms_norm_weight,
564+
epsilon=self.eps,
565+
dst_type=torch.float8_e4m3fn,
566+
)
567+
mxscale = output[2]
568+
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[0], True)
569+
mxscale = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(mxscale, True)
570+
return quantized_output, mxscale, output[1]
571+
572+
return replacement
573+
574+
575+
class RMSNormDynamicMXQuantPattern(BasePattern):
576+
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
577+
super().__init__(vllm_config, eps)
578+
579+
def get_inputs(self):
580+
"""
581+
Generate example inputs for the RMSNormDynamicMXQuant fusion pattern.
582+
"""
583+
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
584+
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
585+
return [rms_norm_input, rms_norm_weight]
586+
587+
def get_pattern(self):
588+
def pattern(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
589+
"""
590+
Pattern for RMSNormDynamicMXQuant fusion.
591+
"""
592+
output = torch.ops.npu.npu_rms_norm(rms_norm_input, rms_norm_weight, self.eps)
593+
out0 = output[0]
594+
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
595+
return quantized_output[0], quantized_output[1]
596+
597+
return pattern
598+
599+
def get_replacement(self):
600+
def replacement(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
601+
"""
602+
Replacement for the RMSNormDynamicMXQuant fusion.
603+
"""
604+
output = torch.ops.npu.npu_rms_norm_dynamic_mx_quant(
605+
rms_norm_input,
606+
rms_norm_weight,
607+
epsilon=self.eps,
608+
dst_type=torch.float8_e4m3fn,
609+
)
610+
return output[0], output[1]
611+
612+
return replacement
613+
614+
615+
class RMSNormDynamicMXQuantSPPattern(BasePattern):
616+
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
617+
super().__init__(vllm_config, eps)
618+
619+
def get_inputs(self):
620+
"""
621+
Generate example inputs for the RMSNormDynamicMXQuant fusion pattern.
622+
"""
623+
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
624+
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
625+
return [rms_norm_input, rms_norm_weight]
626+
627+
def get_pattern(self):
628+
def pattern(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
629+
"""
630+
Pattern for RMSNormDynamicMXQuant fusion.
631+
"""
632+
output = torch.ops.npu.npu_rms_norm(rms_norm_input, rms_norm_weight, self.eps)
633+
out0 = output[0]
634+
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
635+
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
636+
return quantized_output[0], quantized_output[1]
637+
638+
return pattern
639+
640+
def get_replacement(self):
641+
def replacement(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
642+
"""
643+
Replacement for the RMSNormDynamicMXQuant fusion.
644+
"""
645+
output = torch.ops.npu.npu_rms_norm_dynamic_mx_quant(
646+
rms_norm_input,
647+
rms_norm_weight,
648+
epsilon=self.eps,
649+
dst_type=torch.float8_e4m3fn,
650+
)
651+
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[0], True)
652+
mxscale = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[1], True)
653+
return quantized_output, mxscale
654+
655+
return replacement
656+
657+
477658
class AddRMSNormQuantFusionPass(VllmInductorPass):
478659
"""
479660
A pass for fusing AddRMSNorm and W8A8 quantization operations on Ascend.
@@ -488,10 +669,29 @@ def __init__(self, vllm_config: VllmConfig):
488669
logger.debug("Quant fusion not enabled: unsupported dtype %s", dtype)
489670
return
490671

672+
dynamic_mx_quant_fusion_available = is_add_rms_norm_dynamic_mx_quant_fusion_available()
673+
if not dynamic_mx_quant_fusion_available:
674+
logger.debug(
675+
"AddRMSNormDynamicMXQuant fusion not enabled: required MX symbols unavailable, or device isn't A5"
676+
)
677+
678+
rms_norm_dynamic_mx_quant_fusion_available = is_rms_norm_dynamic_mx_quant_fusion_available()
679+
if not rms_norm_dynamic_mx_quant_fusion_available:
680+
logger.debug(
681+
"RMSNormDynamicMXQuant fusion not enabled: required MX symbols unavailable, or device isn't A5"
682+
)
683+
491684
common_epsilons = [1e-5, 1e-6]
685+
492686
for eps in common_epsilons:
493687
AddRMSNormDynamicQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
494688
AddRMSNormDynamicQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
689+
if dynamic_mx_quant_fusion_available:
690+
AddRMSNormDynamicMXQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
691+
AddRMSNormDynamicMXQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
692+
if rms_norm_dynamic_mx_quant_fusion_available:
693+
RMSNormDynamicMXQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
694+
RMSNormDynamicMXQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
495695
if enable_custom_op():
496696
AddRMSNormQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
497697
AddRMSNormQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)

vllm_ascend/device/mxfp_compat.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
import torch
22
import torch_npu
33

4+
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
5+
46
# TODO(linfeng): Temporary compatibility shim for MXFP4/MXFP8 because current torch_npu
57
# releases do not expose the required dtype attributes yet. Simplify or remove this
68
# file after the torch_npu release in March 2026 includes those dtype symbols.
@@ -20,6 +22,10 @@ def _get_missing_symbols(symbols: tuple[str, ...]) -> list[str]:
2022
return [symbol for symbol in symbols if not hasattr(torch_npu, symbol)]
2123

2224

25+
def _is_dynamic_mx_quant_fusion_soc_supported() -> bool:
26+
return get_ascend_device_type() == AscendDeviceType.A5
27+
28+
2329
def _ensure_symbols_available(feature: str, symbols: tuple[str, ...]) -> None:
2430
missing_symbols = _get_missing_symbols(symbols)
2531
if not missing_symbols:
@@ -31,6 +37,22 @@ def _ensure_symbols_available(feature: str, symbols: tuple[str, ...]) -> None:
3137
)
3238

3339

40+
def is_add_rms_norm_dynamic_mx_quant_fusion_available() -> bool:
41+
return (
42+
_is_dynamic_mx_quant_fusion_soc_supported()
43+
and hasattr(torch, "float8_e4m3fn")
44+
and not _get_missing_symbols(("npu_dynamic_mx_quant", "npu_add_rms_norm_dynamic_mx_quant"))
45+
)
46+
47+
48+
def is_rms_norm_dynamic_mx_quant_fusion_available() -> bool:
49+
return (
50+
_is_dynamic_mx_quant_fusion_soc_supported()
51+
and hasattr(torch, "float8_e4m3fn")
52+
and not _get_missing_symbols(("npu_dynamic_mx_quant", "npu_rms_norm_dynamic_mx_quant"))
53+
)
54+
55+
3456
def ensure_mxfp8_scale_dtype_available(feature: str) -> None:
3557
_ensure_symbols_available(feature, ("float8_e8m0fnu",))
3658

0 commit comments

Comments
 (0)