Skip to content

Commit 817ca86

Browse files
Merge branch 'main' into test-cpu-32-hk-verify
2 parents db36970 + bb2666d commit 817ca86

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)