2323from vllm .logger import logger
2424
2525from 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+ )
2630from 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+
477658class 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 )
0 commit comments