0
vllm_ascend/compilation/passes/__init__.py
Normal file
0
vllm_ascend/compilation/passes/__init__.py
Normal file
40
vllm_ascend/compilation/passes/allgather_chunk_noop_pass.py
Normal file
40
vllm_ascend/compilation/passes/allgather_chunk_noop_pass.py
Normal file
@@ -0,0 +1,40 @@
|
||||
import torch
|
||||
import torch._inductor.pattern_matcher as pm
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size, get_tp_group
|
||||
from vllm.logger import logger
|
||||
|
||||
|
||||
class AllGatherChunkNoOpCleanupPass(VllmInductorPass):
|
||||
"""Fold all_gather + sequence_parallel_chunk_impl into identity."""
|
||||
|
||||
def __init__(self, config: VllmConfig):
|
||||
super().__init__(config)
|
||||
self.tp_group = get_tp_group()
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.patterns: PatternMatcherPass = PatternMatcherPass(pass_name="npu_allgather_chunk_noop_cleanup_pass")
|
||||
self._register_patterns()
|
||||
|
||||
def _all_gather(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return torch.ops.vllm.all_gather(x, dim=0, world_size=self.tp_size, group_name=self.tp_group.unique_name)
|
||||
|
||||
def _empty(self, *args, **kwargs):
|
||||
return torch.empty(*args, dtype=self.model_dtype, device=self.device, **kwargs)
|
||||
|
||||
def _register_patterns(self) -> None:
|
||||
def pattern(input: torch.Tensor) -> torch.Tensor:
|
||||
gathered = self._all_gather(input)
|
||||
return torch.ops.vllm.sequence_parallel_chunk_impl(gathered)
|
||||
|
||||
def replacement(input: torch.Tensor) -> torch.Tensor:
|
||||
return input
|
||||
|
||||
pm.register_replacement(pattern, replacement, [self._empty(8, 16)], pm.fwd_only, self.patterns)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph) -> None:
|
||||
self.begin()
|
||||
matched_count = self.patterns.apply(graph)
|
||||
logger.debug("AllGatherChunkNoOpCleanupPass replaced %s patterns", matched_count)
|
||||
self.end_and_log()
|
||||
159
vllm_ascend/compilation/passes/allreduce_rmsnorm_fusion_pass.py
Normal file
159
vllm_ascend/compilation/passes/allreduce_rmsnorm_fusion_pass.py
Normal file
@@ -0,0 +1,159 @@
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
import torch
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass, PatternPrettyPrinter
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.compilation import Range
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size, tensor_model_parallel_all_reduce
|
||||
from vllm.distributed.parallel_state import get_tp_group
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.compilation.passes.base_pattern import BasePattern
|
||||
|
||||
# computation-communication tiling block is 512
|
||||
ALLREDUCE_NORM_FUSE_THRESHOLD = 512
|
||||
|
||||
|
||||
class MiddleLayerMatmulAllReduceAddRMSNormPattern(BasePattern):
|
||||
"""
|
||||
recognizing the Matmul+AllReduce+AddRMSNorm computation pattern
|
||||
AllReduce is optimized in the fusion operator to a two-stage communication of ReduceScatter+AllGather
|
||||
"""
|
||||
|
||||
def __init__(self, vllm_config, eps=1e-6):
|
||||
self.vllm_config = vllm_config
|
||||
self.eps = eps
|
||||
device_group = get_tp_group().device_group
|
||||
backend = device_group._get_backend(torch.device("npu"))
|
||||
self.local_rank = torch.distributed.get_rank(group=device_group)
|
||||
self.tp_group_name = backend.get_hccl_comm_name(self.local_rank)
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
|
||||
def get_inputs(self):
|
||||
batch_size, seq_len = 2, 4
|
||||
hidden_size = 4096
|
||||
x = torch.randn(batch_size, seq_len, hidden_size, device="npu")
|
||||
weight = torch.randn(hidden_size, hidden_size, device="npu")
|
||||
residual = torch.randn(batch_size, seq_len, hidden_size, device="npu")
|
||||
rms_norm_weight = torch.randn(hidden_size, device="npu")
|
||||
return [x, weight, residual, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(x, weight, residual, rms_norm_weight):
|
||||
mm = torch.ops.vllm.unquantized_gemm(x, weight, None)
|
||||
all_reduce_ = tensor_model_parallel_all_reduce(mm)
|
||||
chunked_residual = torch.ops.vllm.maybe_chunk_residual(all_reduce_, residual)
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(all_reduce_, chunked_residual, rms_norm_weight, None)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
return out0, out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(x, weight, residual, rms_norm_weight):
|
||||
out0, out1 = torch.ops._C_ascend.matmul_allreduce_add_rmsnorm(
|
||||
x,
|
||||
weight,
|
||||
residual,
|
||||
rms_norm_weight,
|
||||
self.tp_group_name,
|
||||
self.tp_size,
|
||||
self.local_rank,
|
||||
self.eps,
|
||||
True,
|
||||
False,
|
||||
)
|
||||
return out0, out1
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class LastLayerMatmulAllReduceAddRMSNormPattern(BasePattern):
|
||||
def __init__(self, vllm_config, eps=1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
device_group = get_tp_group().device_group
|
||||
backend = device_group._get_backend(torch.device("npu"))
|
||||
self.local_rank = torch.distributed.get_rank(group=device_group)
|
||||
self.tp_group_name = backend.get_hccl_comm_name(self.local_rank)
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
|
||||
def get_inputs(self):
|
||||
batch_size, seq_len = 2, 4
|
||||
hidden_size = 4096
|
||||
x = torch.randn(batch_size, seq_len, hidden_size, device="npu")
|
||||
weight = torch.randn(hidden_size, hidden_size, device="npu")
|
||||
residual = torch.randn(batch_size, seq_len, hidden_size, device="npu")
|
||||
rms_norm_weight = torch.randn(hidden_size, device="npu")
|
||||
return [x, weight, residual, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(x, weight, residual, rms_norm_weight):
|
||||
mm = torch.ops.vllm.unquantized_gemm(x, weight, None)
|
||||
all_reduce_ = tensor_model_parallel_all_reduce(mm)
|
||||
chunked_residual = torch.ops.vllm.maybe_chunk_residual(all_reduce_, residual)
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(all_reduce_, chunked_residual, rms_norm_weight, None)
|
||||
return output[0]
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(x, weight, residual, rms_norm_weight):
|
||||
out0, _ = torch.ops._C_ascend.matmul_allreduce_add_rmsnorm(
|
||||
x,
|
||||
weight,
|
||||
residual,
|
||||
rms_norm_weight,
|
||||
self.tp_group_name,
|
||||
self.tp_size,
|
||||
self.local_rank,
|
||||
self.eps,
|
||||
True,
|
||||
False,
|
||||
)
|
||||
return out0
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class MatmulAllReduceAddRMSNormPass(VllmInductorPass):
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
super().__init__(vllm_config)
|
||||
self.pattern_match_passes: PatternMatcherPass = PatternMatcherPass(pass_name="allreduce_rmsnorm_fusion_pass")
|
||||
|
||||
MiddleLayerMatmulAllReduceAddRMSNormPattern(vllm_config).register(self.pattern_match_passes)
|
||||
LastLayerMatmulAllReduceAddRMSNormPattern(vllm_config).register(self.pattern_match_passes)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph):
|
||||
self.begin()
|
||||
self.matched_count = self.pattern_match_passes.apply(graph)
|
||||
pattern_idx = 0
|
||||
for pattern_entry in self.pattern_match_passes.patterns.values():
|
||||
for p in pattern_entry:
|
||||
p_str = PatternPrettyPrinter.run(p.pattern)
|
||||
logger.debug("Pattern %d: %s", pattern_idx, p_str)
|
||||
pattern_idx += 1
|
||||
logger.debug("Replaced %s patterns", self.matched_count)
|
||||
self.end_and_log()
|
||||
|
||||
def is_applicable_for_range(self, compile_range: Range) -> bool:
|
||||
"""
|
||||
Check if the pass is applicable for the current configuration.
|
||||
"""
|
||||
applicable = compile_range.start > ALLREDUCE_NORM_FUSE_THRESHOLD
|
||||
return applicable
|
||||
63
vllm_ascend/compilation/passes/base_pattern.py
Normal file
63
vllm_ascend/compilation/passes/base_pattern.py
Normal file
@@ -0,0 +1,63 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
import torch._inductor.pattern_matcher as pm
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
from vllm.config import VllmConfig
|
||||
|
||||
try:
|
||||
import npugraph_ex as nge
|
||||
except ImportError:
|
||||
import torchair as nge
|
||||
|
||||
from vllm_ascend.compilation.passes.utils.npugraph_ex_utils_check import extra_stream_scope_check
|
||||
|
||||
# Global set to track registered patterns and prevent duplicates
|
||||
_registered_patterns: set[str] = set()
|
||||
|
||||
|
||||
class BasePattern(ABC):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
self.vllm_config = vllm_config
|
||||
self.dtype = vllm_config.model_config.dtype
|
||||
self.eps = eps
|
||||
|
||||
@abstractmethod
|
||||
def get_inputs(self) -> list[torch.Tensor]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_pattern(self) -> Callable:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_replacement(self) -> Callable:
|
||||
pass
|
||||
|
||||
def get_extra_stream_scope_check(self):
|
||||
return extra_stream_scope_check
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass) -> None:
|
||||
# Create a unique identifier for this pattern based on class name and eps
|
||||
pattern_id = f"{self.__class__.__name__}_{self.eps}"
|
||||
|
||||
# Skip registration if this pattern has already been registered globally
|
||||
if pattern_id in _registered_patterns:
|
||||
return
|
||||
|
||||
pattern_fn = self.get_pattern()
|
||||
replacement_fn = self.get_replacement()
|
||||
example_inputs = self.get_inputs()
|
||||
|
||||
pm.register_replacement(pattern_fn, replacement_fn, example_inputs, pm.fwd_only, pm_pass)
|
||||
|
||||
nge.register_replacement(
|
||||
search_fn=pattern_fn,
|
||||
replace_fn=replacement_fn,
|
||||
example_inputs=example_inputs,
|
||||
extra_check=self.get_extra_stream_scope_check(),
|
||||
)
|
||||
|
||||
# Mark this pattern as registered
|
||||
_registered_patterns.add(pattern_id)
|
||||
110
vllm_ascend/compilation/passes/muls_add_pass.py
Normal file
110
vllm_ascend/compilation/passes/muls_add_pass.py
Normal file
@@ -0,0 +1,110 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.compilation import Range
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.compilation.passes.base_pattern import BasePattern
|
||||
|
||||
|
||||
class MulsAddPattern(BasePattern):
|
||||
"""
|
||||
Pattern that matches an element-wise mul + add sequence:
|
||||
tmp = x * scale
|
||||
out = tmp + y
|
||||
and replaces it with a call to the muls_add_triton kernel.
|
||||
"""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, scale: float = 1.0):
|
||||
super().__init__(vllm_config)
|
||||
self.scale = scale
|
||||
|
||||
def get_inputs(self) -> list[torch.Tensor]:
|
||||
"""
|
||||
Generate example inputs for the MulsAddPattern.
|
||||
|
||||
The exact shapes are not important for pattern matching; they only
|
||||
provide meta information for the pattern matcher.
|
||||
"""
|
||||
x = torch.randn(2, 2048, device="npu", dtype=self.dtype)
|
||||
y = torch.randn(2, 2048, device="npu", dtype=self.dtype)
|
||||
# Only tensor inputs are needed here. The scalar scale is stored on the
|
||||
# pattern instance (self.scale) instead of being passed as an input.
|
||||
return [x, y]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(x: torch.Tensor, y: torch.Tensor):
|
||||
"""
|
||||
Pattern for element-wise x * scale + y.
|
||||
"""
|
||||
tmp = x * self.scale
|
||||
out = tmp + y
|
||||
return out
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(x: torch.Tensor, y: torch.Tensor):
|
||||
"""
|
||||
Replacement that calls the muls_add_triton kernel using the
|
||||
class-level scalar self.scale.
|
||||
"""
|
||||
return torch.ops.vllm.muls_add(x, y, self.scale)
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class MulsAddFusionPass(VllmInductorPass):
|
||||
"""
|
||||
A fusion pass that replaces simple element-wise x * scale + y patterns
|
||||
with the Triton-based muls_add_triton kernel on Ascend.
|
||||
"""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
super().__init__(vllm_config)
|
||||
self.pattern_match_passes: PatternMatcherPass = PatternMatcherPass(pass_name="muls_add_fusion_pass")
|
||||
|
||||
# For now we enable this pass for all floating-point dtypes that the
|
||||
# model is configured to use.
|
||||
dtype = vllm_config.model_config.dtype
|
||||
if dtype not in (torch.float16, torch.bfloat16, torch.float32):
|
||||
logger.debug("MulsAdd fusion not enabled: unsupported dtype %s", dtype)
|
||||
return
|
||||
|
||||
routed_scaling_factor = getattr(vllm_config.model_config.hf_text_config, "routed_scaling_factor", 1.0)
|
||||
MulsAddPattern(vllm_config, scale=routed_scaling_factor).register(self.pattern_match_passes)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph) -> None: # type: ignore[override]
|
||||
self.begin()
|
||||
self.matched_count = self.pattern_match_passes.apply(graph)
|
||||
logger.debug("Fused %s muls_add patterns", self.matched_count)
|
||||
self.end_and_log()
|
||||
|
||||
def is_applicable_for_range(self, compile_range: Range) -> bool:
|
||||
"""
|
||||
Check if the pass is applicable for the current configuration.
|
||||
|
||||
For now, muls_add fusion is always allowed for the selected ranges.
|
||||
This hook exists so that we can add more fine-grained range control
|
||||
in the future if needed.
|
||||
"""
|
||||
return True
|
||||
62
vllm_ascend/compilation/passes/noop_elimination.py
Normal file
62
vllm_ascend/compilation/passes/noop_elimination.py
Normal file
@@ -0,0 +1,62 @@
|
||||
from collections.abc import Iterable
|
||||
|
||||
import torch
|
||||
import torch.fx
|
||||
from torch import SymInt
|
||||
from torch.fx.experimental.symbolic_shapes import statically_known_true
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.logger import logger
|
||||
|
||||
|
||||
class NoOpEliminationPass(VllmInductorPass):
|
||||
"""Remove no-op view/reshape nodes after pattern rewrites."""
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph) -> None:
|
||||
fx_graph = graph.graph if hasattr(graph, "graph") else graph
|
||||
removed = 0
|
||||
for node in list(fx_graph.nodes):
|
||||
if not self._is_view_like(node):
|
||||
continue
|
||||
|
||||
input_node = node.args[0]
|
||||
if not isinstance(input_node, torch.fx.Node):
|
||||
continue
|
||||
|
||||
input_meta = input_node.meta.get("val")
|
||||
output_meta = node.meta.get("val")
|
||||
if input_meta is None or output_meta is None:
|
||||
continue
|
||||
|
||||
input_shape = getattr(input_meta, "shape", None)
|
||||
output_shape = getattr(output_meta, "shape", None)
|
||||
if input_shape is None or output_shape is None:
|
||||
continue
|
||||
|
||||
if self._all_dims_equivalent(input_shape, output_shape):
|
||||
node.replace_all_uses_with(input_node)
|
||||
fx_graph.erase_node(node)
|
||||
removed += 1
|
||||
|
||||
logger.debug("NoOpEliminationPass removed %s no-op views", removed)
|
||||
|
||||
@staticmethod
|
||||
def _is_view_like(node: torch.fx.Node) -> bool:
|
||||
return (node.op == "call_method" and node.target in {"view", "reshape"}) or (
|
||||
node.op == "call_function"
|
||||
and node.target
|
||||
in {
|
||||
torch.ops.aten.view.default,
|
||||
torch.ops.aten.reshape.default,
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _dims_equivalent(dim: int | SymInt, i_dim: int | SymInt) -> bool:
|
||||
return statically_known_true(dim == i_dim) # type: ignore[no-any-return]
|
||||
|
||||
def _all_dims_equivalent(self, dims: Iterable[int | SymInt], i_dims: Iterable[int | SymInt]) -> bool:
|
||||
dims_ = list(dims)
|
||||
i_dims_ = list(i_dims)
|
||||
if len(dims_) != len(i_dims_):
|
||||
return False
|
||||
return all(self._dims_equivalent(s, i_s) for s, i_s in zip(dims_, i_dims_))
|
||||
742
vllm_ascend/compilation/passes/norm_quant_fusion_pass.py
Normal file
742
vllm_ascend/compilation/passes/norm_quant_fusion_pass.py
Normal file
@@ -0,0 +1,742 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
import torch
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.compilation import Range
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.compilation.passes.base_pattern import BasePattern
|
||||
from vllm_ascend.device.mxfp_compat import (
|
||||
is_add_rms_norm_dynamic_mx_quant_fusion_available,
|
||||
is_rms_norm_dynamic_mx_quant_fusion_available,
|
||||
)
|
||||
from vllm_ascend.utils import enable_custom_op
|
||||
|
||||
|
||||
class AddRMSNormQuantPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
scale = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
scale_reciprocal = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
offset = torch.zeros(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight, scale, scale_reciprocal, offset]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
rms_norm_input, residual, rms_norm_weight, None, self.eps
|
||||
)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.vllm.quantize(out0, scale, scale_reciprocal, offset)
|
||||
return quantized_output, out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, scale, offset, epsilon=self.eps
|
||||
)
|
||||
quantized_output = output[0]
|
||||
out1 = output[2]
|
||||
return quantized_output, out1
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormQuantPatternWithBias(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
rmsnorm_bias = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
scale = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
scale_reciprocal = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
offset = torch.zeros(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight, scale, scale_reciprocal, offset, rmsnorm_bias]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
rms_norm_input, residual, rms_norm_weight, bias, self.eps
|
||||
)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.vllm.quantize(out0, scale, scale_reciprocal, offset)
|
||||
return quantized_output, out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, scale, offset, epsilon=self.eps, beta=bias
|
||||
)
|
||||
quantized_output = output[0]
|
||||
out1 = output[2]
|
||||
return quantized_output, out1
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormQuantSPPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
scale = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
scale_reciprocal = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
offset = torch.zeros(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight, scale, scale_reciprocal, offset]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
rms_norm_input, residual, rms_norm_weight, None, self.eps
|
||||
)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
|
||||
quantized_output = torch.ops.vllm.quantize(out0, scale, scale_reciprocal, offset)
|
||||
return quantized_output, out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, scale, offset, epsilon=self.eps
|
||||
)
|
||||
quantized_output = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(quantized_output, True)
|
||||
return quantized_output, out1
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormQuantSPPatternWithBias(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
rmsnorm_bias = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
scale = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
scale_reciprocal = torch.ones(4, device="npu", dtype=self.dtype)
|
||||
offset = torch.zeros(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight, scale, scale_reciprocal, offset, rmsnorm_bias]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
rms_norm_input, residual, rms_norm_weight, bias, self.eps
|
||||
)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
|
||||
quantized_output = torch.ops.vllm.quantize(out0, scale, scale_reciprocal, offset)
|
||||
return quantized_output, out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
scale_reciprocal: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, scale, offset, epsilon=self.eps, beta=bias
|
||||
)
|
||||
quantized_output = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(quantized_output, True)
|
||||
return quantized_output, out1
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormDynamicQuantPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm(rms_norm_input, residual, rms_norm_weight, self.eps)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.npu.npu_dynamic_quant(out0)
|
||||
return quantized_output[0], quantized_output[1], out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_dynamic_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, epsilon=self.eps, output_mask=[True, False]
|
||||
)
|
||||
return (
|
||||
output[0],
|
||||
output[3],
|
||||
output[2],
|
||||
)
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormDynamicQuantPatternWithBias(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
rmsnorm_bias = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight, rmsnorm_bias]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
rms_norm_input, residual, rms_norm_weight, bias, self.eps
|
||||
)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.npu.npu_dynamic_quant(out0)
|
||||
return quantized_output[0], quantized_output[1], out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_dynamic_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, epsilon=self.eps, output_mask=[True, False], beta=bias
|
||||
)
|
||||
return (
|
||||
output[0],
|
||||
output[3],
|
||||
output[2],
|
||||
)
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormDynamicQuantSPPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm(rms_norm_input, residual, rms_norm_weight, self.eps)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
|
||||
quantized_output = torch.ops.npu.npu_dynamic_quant(out0)
|
||||
return quantized_output[0], quantized_output[1], out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_dynamic_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, epsilon=self.eps, output_mask=[True, False]
|
||||
)
|
||||
out3 = output[3]
|
||||
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[0], True)
|
||||
out3 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out3, True)
|
||||
return quantized_output, out3, output[2]
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormDynamicQuantSPPatternWithBias(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 4, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
rmsnorm_bias = torch.randn(4, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight, rmsnorm_bias]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Pattern for AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
rms_norm_input, residual, rms_norm_weight, bias, self.eps
|
||||
)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
|
||||
quantized_output = torch.ops.npu.npu_dynamic_quant(out0)
|
||||
return quantized_output[0], quantized_output[1], out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
rms_norm_input: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
rms_norm_weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Replacement for the AddRMSNormQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_dynamic_quant(
|
||||
rms_norm_input, residual, rms_norm_weight, epsilon=self.eps, output_mask=[True, False], beta=bias
|
||||
)
|
||||
out3 = output[3]
|
||||
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[0], True)
|
||||
out3 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out3, True)
|
||||
return quantized_output, out3, output[2]
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormDynamicMXQuantPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormDynamicMXQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 64, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Pattern for AddRMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm(rms_norm_input, residual, rms_norm_weight, self.eps)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
|
||||
return quantized_output[0], quantized_output[1], out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Replacement for the AddRMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_dynamic_mx_quant(
|
||||
rms_norm_input,
|
||||
residual,
|
||||
rms_norm_weight,
|
||||
epsilon=self.eps,
|
||||
dst_type=torch.float8_e4m3fn,
|
||||
)
|
||||
return (
|
||||
output[0],
|
||||
output[2],
|
||||
output[1],
|
||||
)
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class AddRMSNormDynamicMXQuantSPPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the AddRMSNormDynamicMXQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
|
||||
residual = torch.randn(2, 64, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, residual, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Pattern for AddRMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm(rms_norm_input, residual, rms_norm_weight, self.eps)
|
||||
out0 = output[0]
|
||||
out1 = output[2]
|
||||
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
|
||||
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
|
||||
return quantized_output[0], quantized_output[1], out1
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(rms_norm_input: torch.Tensor, residual: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Replacement for the AddRMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_add_rms_norm_dynamic_mx_quant(
|
||||
rms_norm_input,
|
||||
residual,
|
||||
rms_norm_weight,
|
||||
epsilon=self.eps,
|
||||
dst_type=torch.float8_e4m3fn,
|
||||
)
|
||||
mxscale = output[2]
|
||||
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[0], True)
|
||||
mxscale = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(mxscale, True)
|
||||
return quantized_output, mxscale, output[1]
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class RMSNormDynamicMXQuantPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the RMSNormDynamicMXQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Pattern for RMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_rms_norm(rms_norm_input, rms_norm_weight, self.eps)
|
||||
out0 = output[0]
|
||||
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
|
||||
return quantized_output[0], quantized_output[1]
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Replacement for the RMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_rms_norm_dynamic_mx_quant(
|
||||
rms_norm_input,
|
||||
rms_norm_weight,
|
||||
epsilon=self.eps,
|
||||
dst_type=torch.float8_e4m3fn,
|
||||
)
|
||||
return output[0], output[1]
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class RMSNormDynamicMXQuantSPPattern(BasePattern):
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs for the RMSNormDynamicMXQuant fusion pattern.
|
||||
"""
|
||||
rms_norm_input = torch.randn(2, 64, device="npu", dtype=self.dtype)
|
||||
rms_norm_weight = torch.randn(64, device="npu", dtype=self.dtype)
|
||||
return [rms_norm_input, rms_norm_weight]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Pattern for RMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_rms_norm(rms_norm_input, rms_norm_weight, self.eps)
|
||||
out0 = output[0]
|
||||
out0 = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(out0, True)
|
||||
quantized_output = torch.ops.npu.npu_dynamic_mx_quant(out0, dst_type=torch.float8_e4m3fn)
|
||||
return quantized_output[0], quantized_output[1]
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(rms_norm_input: torch.Tensor, rms_norm_weight: torch.Tensor):
|
||||
"""
|
||||
Replacement for the RMSNormDynamicMXQuant fusion.
|
||||
"""
|
||||
output = torch.ops.npu.npu_rms_norm_dynamic_mx_quant(
|
||||
rms_norm_input,
|
||||
rms_norm_weight,
|
||||
epsilon=self.eps,
|
||||
dst_type=torch.float8_e4m3fn,
|
||||
)
|
||||
quantized_output = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[0], True)
|
||||
mxscale = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(output[1], True)
|
||||
return quantized_output, mxscale
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
def _model_uses_w4a4_quant(vllm_config: VllmConfig | None) -> bool:
|
||||
"""Check whether the model uses W4A4 int4 quantization for any layer.
|
||||
|
||||
W4A4 int4 schemes (e.g. W4A4_DYNAMIC, W4A4_FLATQUANT_DYNAMIC)
|
||||
are incompatible with the fuse_norm_quant optimization,
|
||||
so callers use this to disable that fusion.
|
||||
"""
|
||||
if vllm_config is None:
|
||||
return False
|
||||
quant_config = getattr(vllm_config, "quant_config", None)
|
||||
if quant_config is None:
|
||||
return False
|
||||
quant_description = getattr(quant_config, "quant_description", None)
|
||||
if not quant_description:
|
||||
return False
|
||||
w4a4_int4_schemes = ["W4A4_DYNAMIC", "W4A4_FLATQUANT_DYNAMIC"]
|
||||
return any(
|
||||
isinstance(quant_type, str) and quant_type in w4a4_int4_schemes for quant_type in quant_description.values()
|
||||
)
|
||||
|
||||
|
||||
class AddRMSNormQuantFusionPass(VllmInductorPass):
|
||||
"""
|
||||
A pass for fusing AddRMSNorm and W8A8 quantization operations on Ascend.
|
||||
"""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
super().__init__(vllm_config)
|
||||
self.pattern_match_passes: PatternMatcherPass = PatternMatcherPass(pass_name="rmsnorm_quant_fusion_pass")
|
||||
|
||||
dtype = vllm_config.model_config.dtype
|
||||
if dtype not in (torch.bfloat16, torch.float16):
|
||||
logger.debug("Quant fusion not enabled: unsupported dtype %s", dtype)
|
||||
return
|
||||
|
||||
if _model_uses_w4a4_quant(vllm_config):
|
||||
logger.debug(
|
||||
"Quant fusion not enabled: the model contains "
|
||||
"W4A4 quantized weights, which are incompatible with the "
|
||||
"norm-quant fusion pass."
|
||||
)
|
||||
return
|
||||
|
||||
dynamic_mx_quant_fusion_available = is_add_rms_norm_dynamic_mx_quant_fusion_available()
|
||||
if not dynamic_mx_quant_fusion_available:
|
||||
logger.debug(
|
||||
"AddRMSNormDynamicMXQuant fusion not enabled: required MX symbols unavailable, or device isn't A5"
|
||||
)
|
||||
|
||||
rms_norm_dynamic_mx_quant_fusion_available = is_rms_norm_dynamic_mx_quant_fusion_available()
|
||||
if not rms_norm_dynamic_mx_quant_fusion_available:
|
||||
logger.debug(
|
||||
"RMSNormDynamicMXQuant fusion not enabled: required MX symbols unavailable, or device isn't A5"
|
||||
)
|
||||
|
||||
common_epsilons = [1e-5, 1e-6]
|
||||
|
||||
for eps in common_epsilons:
|
||||
AddRMSNormDynamicQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormDynamicQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
if dynamic_mx_quant_fusion_available:
|
||||
AddRMSNormDynamicMXQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormDynamicMXQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
if rms_norm_dynamic_mx_quant_fusion_available:
|
||||
RMSNormDynamicMXQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
RMSNormDynamicMXQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
if enable_custom_op():
|
||||
AddRMSNormQuantPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormQuantSPPattern(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormQuantPatternWithBias(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormQuantSPPatternWithBias(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormDynamicQuantPatternWithBias(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
AddRMSNormDynamicQuantSPPatternWithBias(vllm_config, eps=eps).register(self.pattern_match_passes)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph):
|
||||
self.begin()
|
||||
self.matched_count = self.pattern_match_passes.apply(graph)
|
||||
logger.debug("Replaced %s patterns", self.matched_count)
|
||||
self.end_and_log()
|
||||
|
||||
def is_applicable_for_range(self, compile_range: Range) -> bool:
|
||||
"""
|
||||
Check if the pass is applicable for the current configuration.
|
||||
"""
|
||||
return True
|
||||
244
vllm_ascend/compilation/passes/qknorm_rope_fusion_pass.py
Normal file
244
vllm_ascend/compilation/passes/qknorm_rope_fusion_pass.py
Normal file
@@ -0,0 +1,244 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
import torch
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass, PatternPrettyPrinter
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.config import VllmConfig, get_layers_from_vllm_config
|
||||
from vllm.config.compilation import Range
|
||||
from vllm.logger import logger
|
||||
from vllm.model_executor.layers.attention import Attention
|
||||
|
||||
from vllm_ascend.compilation.passes.base_pattern import BasePattern
|
||||
from vllm_ascend.device.device_op import DeviceOperator
|
||||
from vllm_ascend.utils import get_rope_dim
|
||||
|
||||
|
||||
class QKNormRopeFusionPattern(BasePattern):
|
||||
def __init__(self, vllm_config, head_dim, num_heads, num_kv_heads, eps=1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
self.head_dim = head_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.device = vllm_config.device_config.device if vllm_config.device_config else None
|
||||
self.rope_dim = get_rope_dim(vllm_config)
|
||||
|
||||
def get_inputs(self):
|
||||
T = 5
|
||||
max_position_embeddings = 16384
|
||||
qkv = torch.empty(T, self.q_size + 2 * self.kv_size, dtype=torch.bfloat16, device="npu")
|
||||
q_weight = torch.empty(self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
k_weight = torch.empty(self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
cos_sin_cache = torch.empty(max_position_embeddings, self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
positions = torch.ones(T, dtype=torch.int64, device="npu")
|
||||
return [qkv, q_weight, k_weight, cos_sin_cache, positions]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
qkv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
):
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
|
||||
q_by_head = q.view(*q.shape[:-1], q.shape[-1] // self.head_dim, self.head_dim)
|
||||
q_norm_out, _ = torch.ops.npu.npu_rms_norm(q_by_head, q_weight, self.eps)
|
||||
|
||||
k_by_head = k.view(*k.shape[:-1], k.shape[-1] // self.head_dim, self.head_dim)
|
||||
k_norm_out, _ = torch.ops.npu.npu_rms_norm(k_by_head, k_weight, self.eps)
|
||||
|
||||
q_flat = q_norm_out.view(q.shape)
|
||||
k_flat = k_norm_out.view(k.shape)
|
||||
q_rope, k_rope = torch.ops.vllm.npu_rotary_embedding(
|
||||
positions, q_flat, k_flat, cos_sin_cache, self.head_dim, self.rope_dim, True
|
||||
)
|
||||
|
||||
return q_rope, k_rope, v
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
qkv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
):
|
||||
results = DeviceOperator.split_qkv_rmsnorm_rope(
|
||||
input=qkv,
|
||||
q_weight=q_weight,
|
||||
k_weight=k_weight,
|
||||
q_hidden_size=self.q_size,
|
||||
kv_hidden_size=self.kv_size,
|
||||
head_dim=self.head_dim,
|
||||
eps=self.eps,
|
||||
q_bias=None,
|
||||
k_bias=None,
|
||||
cos_sin_cache=cos_sin_cache,
|
||||
positions=positions,
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class QKNormRopeFusionPatternWithBias(BasePattern):
|
||||
def __init__(self, vllm_config, head_dim, num_heads, num_kv_heads, eps=1e-6):
|
||||
super().__init__(vllm_config, eps)
|
||||
self.head_dim = head_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.device = vllm_config.device_config.device if vllm_config.device_config else None
|
||||
self.rope_dim = get_rope_dim(vllm_config)
|
||||
|
||||
def get_inputs(self):
|
||||
T = 5
|
||||
max_position_embeddings = 16384
|
||||
qkv = torch.empty(T, self.q_size + 2 * self.kv_size, dtype=torch.bfloat16, device="npu")
|
||||
q_weight = torch.empty(self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
k_weight = torch.empty(self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
q_bias = torch.empty(self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
k_bias = torch.empty(self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
cos_sin_cache = torch.empty(max_position_embeddings, self.head_dim, dtype=torch.bfloat16, device="npu")
|
||||
positions = torch.ones(T, dtype=torch.int64, device="npu")
|
||||
|
||||
return [qkv, q_weight, k_weight, q_bias, k_bias, cos_sin_cache, positions]
|
||||
|
||||
def get_pattern(self):
|
||||
def pattern(
|
||||
qkv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_bias: torch.Tensor,
|
||||
k_bias: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
):
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
|
||||
q_by_head = q.view(*q.shape[:-1], q.shape[-1] // self.head_dim, self.head_dim)
|
||||
q_norm_out, _ = torch.ops.npu.npu_rms_norm(q_by_head, q_weight, self.eps)
|
||||
q_normed = q_norm_out + q_bias
|
||||
|
||||
k_by_head = k.view(*k.shape[:-1], k.shape[-1] // self.head_dim, self.head_dim)
|
||||
k_norm_out, _ = torch.ops.npu.npu_rms_norm(k_by_head, k_weight, self.eps)
|
||||
k_normed = k_norm_out + k_bias
|
||||
|
||||
q_flat = q_normed.view(q.shape)
|
||||
k_flat = k_normed.view(k.shape)
|
||||
q_rope, k_rope = torch.ops.vllm.npu_rotary_embedding(
|
||||
positions, q_flat, k_flat, cos_sin_cache, self.head_dim, self.rope_dim, True
|
||||
)
|
||||
|
||||
return q_rope, k_rope, v
|
||||
|
||||
return pattern
|
||||
|
||||
def get_replacement(self):
|
||||
def replacement(
|
||||
qkv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
q_bias: torch.Tensor,
|
||||
k_bias: torch.Tensor,
|
||||
cos_sin_cache: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
):
|
||||
results = DeviceOperator.split_qkv_rmsnorm_rope(
|
||||
input=qkv,
|
||||
q_weight=q_weight,
|
||||
k_weight=k_weight,
|
||||
q_hidden_size=self.q_size,
|
||||
kv_hidden_size=self.kv_size,
|
||||
head_dim=self.head_dim,
|
||||
eps=self.eps,
|
||||
q_bias=q_bias,
|
||||
k_bias=k_bias,
|
||||
cos_sin_cache=cos_sin_cache,
|
||||
positions=positions,
|
||||
)
|
||||
return results
|
||||
|
||||
return replacement
|
||||
|
||||
|
||||
class QKNormRopeFusionPass(VllmInductorPass):
|
||||
"""
|
||||
A pass for fusing QKV split and RMSNorm operations into a single qk_rmsnorm operator.
|
||||
"""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig):
|
||||
super().__init__(vllm_config)
|
||||
self.pattern_match_passes: PatternMatcherPass = PatternMatcherPass(pass_name="qknorm_rope_fusion_pass")
|
||||
|
||||
dtype = vllm_config.model_config.dtype
|
||||
if dtype not in (torch.bfloat16,):
|
||||
logger.debug("QKNorm and Rope fusion not enabled: unsupported dtype %s", dtype)
|
||||
return
|
||||
|
||||
# use one attn layer to get meta (such as head_dim) for QKNormRopeFusionPattern
|
||||
attn_layers: dict[str, Attention] = get_layers_from_vllm_config(vllm_config, Attention)
|
||||
if len(attn_layers) == 0:
|
||||
logger.debug("QKNorm and Rope fusion enabled, but no Attention layers were discovered.")
|
||||
return
|
||||
layer = next(iter(attn_layers.values()))
|
||||
for epsilon in [1e-6, 1e-5]:
|
||||
if layer.head_size != 128:
|
||||
logger.debug("QKNorm and Rope fusion not enabled: head_dim %d is not equal of 128", layer.head_size)
|
||||
continue
|
||||
QKNormRopeFusionPattern(
|
||||
vllm_config=vllm_config,
|
||||
head_dim=layer.head_size,
|
||||
num_heads=layer.num_heads,
|
||||
num_kv_heads=layer.num_kv_heads,
|
||||
eps=epsilon,
|
||||
).register(self.pattern_match_passes)
|
||||
|
||||
QKNormRopeFusionPatternWithBias(
|
||||
vllm_config=vllm_config,
|
||||
head_dim=layer.head_size,
|
||||
num_heads=layer.num_heads,
|
||||
num_kv_heads=layer.num_kv_heads,
|
||||
eps=epsilon,
|
||||
).register(self.pattern_match_passes)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph):
|
||||
self.begin()
|
||||
self.matched_count = self.pattern_match_passes.apply(graph)
|
||||
logger.debug("Fused %s QKNorm and Rope patterns", self.matched_count)
|
||||
logger.debug("Patterns registered for replacement:")
|
||||
pattern_idx = 0
|
||||
for pattern_entry in self.pattern_match_passes.patterns.values():
|
||||
for p in pattern_entry:
|
||||
p_str = PatternPrettyPrinter.run(p.pattern)
|
||||
logger.debug("Pattern %d: %s", pattern_idx, p_str)
|
||||
pattern_idx += 1
|
||||
self.end_and_log()
|
||||
|
||||
def is_applicable_for_range(self, compile_range: Range) -> bool:
|
||||
"""
|
||||
Check if the pass is applicable for the current configuration.
|
||||
"""
|
||||
return True
|
||||
234
vllm_ascend/compilation/passes/sequence_parallelism.py
Normal file
234
vllm_ascend/compilation/passes/sequence_parallelism.py
Normal file
@@ -0,0 +1,234 @@
|
||||
import torch
|
||||
import torch._inductor.pattern_matcher as pm
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
from vllm.compilation.passes.vllm_inductor_pass import VllmInductorPass
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.utils import Range
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size, get_tp_group, tensor_model_parallel_all_reduce
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.compilation.passes.noop_elimination import NoOpEliminationPass
|
||||
from vllm_ascend.utils import is_moe_model
|
||||
|
||||
SP_MIN_TOKEN_NUM_DEFAULT = 1000
|
||||
|
||||
|
||||
def get_sp_min_token_num(config: VllmConfig) -> int:
|
||||
if is_moe_model(config):
|
||||
return 1
|
||||
|
||||
return SP_MIN_TOKEN_NUM_DEFAULT
|
||||
|
||||
|
||||
class _SequenceParallelPatternHelper:
|
||||
"""Helper for sequence parallelism patterns.
|
||||
|
||||
Provides TP communication helper methods: _all_reduce, _reduce_scatter,
|
||||
_all_gather, and tensor creation utilities.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
epsilon: float,
|
||||
dtype: torch.dtype,
|
||||
device: str,
|
||||
):
|
||||
self.eps = epsilon
|
||||
self.dtype = dtype
|
||||
self.device = device
|
||||
self.tp_group = get_tp_group()
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.tp_rank = get_tp_group().rank_in_group
|
||||
|
||||
def _all_reduce(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return tensor_model_parallel_all_reduce(x)
|
||||
|
||||
def _reduce_scatter(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return torch.ops.vllm.reduce_scatter(x, dim=0, world_size=self.tp_size, group_name=self.tp_group.unique_name)
|
||||
|
||||
def _all_gather(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return torch.ops.vllm.all_gather(x, dim=0, world_size=self.tp_size, group_name=self.tp_group.unique_name)
|
||||
|
||||
def empty(self, *args, **kws):
|
||||
return torch.empty(*args, dtype=self.dtype, device="npu", **kws)
|
||||
|
||||
|
||||
class MiddleAllReduceRMSNormPattern(_SequenceParallelPatternHelper):
|
||||
"""Replaces all_reduce + AddRMSNormBias with reduce_scatter + AddRMSNormBias
|
||||
+ all_gather for middle-layer sequence parallelism."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def empty(self, *args, **kws):
|
||||
return torch.empty(*args, dtype=self.dtype, device="npu", **kws)
|
||||
|
||||
def get_inputs(self):
|
||||
"""
|
||||
Generate example inputs.
|
||||
"""
|
||||
input = self.empty(8, 16)
|
||||
weight = self.empty(16)
|
||||
residual = self.empty(8, 16)
|
||||
return [input, weight, residual]
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
x = self._all_reduce(input)
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(x, residual, weight, None, self.eps)
|
||||
|
||||
return result, residual
|
||||
|
||||
def replacement(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
reduce_scatter = self._reduce_scatter(input)
|
||||
residual = torch.ops.vllm.maybe_chunk_residual(reduce_scatter, residual)
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(
|
||||
reduce_scatter, residual, weight, None, self.eps
|
||||
)
|
||||
all_gather = self._all_gather(result)
|
||||
return all_gather, residual
|
||||
|
||||
pm.register_replacement(pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass)
|
||||
|
||||
|
||||
class LastAllReduceRMSNormPattern(_SequenceParallelPatternHelper):
|
||||
"""Same as MiddleAllReduceRMSNormPattern but for the last layer
|
||||
(no residual backprop)."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def get_inputs(self):
|
||||
input = self.empty(8, 16)
|
||||
weight = self.empty(16)
|
||||
residual = self.empty(8, 16)
|
||||
return [input, weight, residual]
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
x = self._all_reduce(input)
|
||||
result, _, _ = torch.ops._C_ascend.npu_add_rms_norm_bias(x, residual, weight, None, self.eps)
|
||||
|
||||
return result
|
||||
|
||||
def replacement(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
reduce_scatter = self._reduce_scatter(input)
|
||||
residual = torch.ops.vllm.maybe_chunk_residual(reduce_scatter, residual)
|
||||
result, _, _ = torch.ops._C_ascend.npu_add_rms_norm_bias(reduce_scatter, residual, weight, None, self.eps)
|
||||
all_gather = self._all_gather(result)
|
||||
return all_gather
|
||||
|
||||
pm.register_replacement(pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass)
|
||||
|
||||
|
||||
class Qwen3VLMiddleAllReduceRMSNormPattern(_SequenceParallelPatternHelper):
|
||||
"""For Qwen3-VL middle layers with hidden_states + deepstack_input_embeds add.
|
||||
|
||||
Replaces all_reduce + add + AddRMSNormBias with reduce_scatter +
|
||||
chunk(deepstack_input_embeds) + add + AddRMSNormBias + all_gather.
|
||||
"""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def get_inputs(self):
|
||||
input = self.empty(8, 16)
|
||||
weight = self.empty(16)
|
||||
residual = self.empty(8, 16)
|
||||
deepstack_input_embeds = self.empty(8, 16)
|
||||
return [input, weight, residual, deepstack_input_embeds]
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
deepstack_input_embeds: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
x = self._all_reduce(input)
|
||||
add_ = x + deepstack_input_embeds
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(add_, residual, weight, None, self.eps)
|
||||
|
||||
return result, residual
|
||||
|
||||
def replacement(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
deepstack_input_embeds: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
reduce_scatter = self._reduce_scatter(input)
|
||||
chunk = deepstack_input_embeds.chunk(self.tp_size)[self.tp_rank]
|
||||
add_ = reduce_scatter + chunk
|
||||
residual = torch.ops.vllm.maybe_chunk_residual(reduce_scatter, residual)
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(add_, residual, weight, None, self.eps)
|
||||
all_gather = self._all_gather(result)
|
||||
return all_gather, residual
|
||||
|
||||
pm.register_replacement(pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass)
|
||||
|
||||
|
||||
class SequenceParallelismPass(VllmInductorPass):
|
||||
"""Sequence parallelism compilation pass.
|
||||
|
||||
Registers and applies the above patterns. Runs noop cleanup first, then
|
||||
uses token range to determine whether to enable SP.
|
||||
"""
|
||||
|
||||
def __init__(self, config: VllmConfig):
|
||||
super().__init__(config)
|
||||
|
||||
self.patterns: PatternMatcherPass = PatternMatcherPass(pass_name="npu_sequence_parallelism_pass")
|
||||
self.noop_cleanup = NoOpEliminationPass(config)
|
||||
|
||||
for epsilon in [1e-5, 1e-6]:
|
||||
MiddleAllReduceRMSNormPattern(config, epsilon).register(self.patterns)
|
||||
|
||||
LastAllReduceRMSNormPattern(config, epsilon).register(self.patterns)
|
||||
|
||||
Qwen3VLMiddleAllReduceRMSNormPattern(config, epsilon).register(self.patterns)
|
||||
|
||||
self.min_tokens = get_sp_min_token_num(config)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph):
|
||||
self.begin()
|
||||
self.noop_cleanup(graph) # Eliminate redundant view-like operations
|
||||
logger.debug("after noop_cleanup %s", graph.graph)
|
||||
self.matched_count = self.patterns.apply(graph)
|
||||
logger.debug("Replaced %s patterns", self.matched_count)
|
||||
logger.debug("after apply replacement %s", graph.graph)
|
||||
|
||||
from torch._inductor.pattern_matcher import PatternPrettyPrinter
|
||||
|
||||
pattern_idx = 0
|
||||
for pattern_entry in self.patterns.patterns.values():
|
||||
for p in pattern_entry:
|
||||
p_str = PatternPrettyPrinter.run(p.pattern)
|
||||
logger.debug("Pattern %d: %s", pattern_idx, p_str)
|
||||
pattern_idx += 1
|
||||
|
||||
self.end_and_log()
|
||||
|
||||
def is_applicable_for_range(self, compile_range: Range) -> bool:
|
||||
"""
|
||||
Check if the pass is applicable for the current configuration.
|
||||
"""
|
||||
applicable = compile_range.start >= self.min_tokens
|
||||
logger.debug("SequenceParallelismPass compile_range=%r applicable=%r", compile_range, applicable)
|
||||
return applicable
|
||||
204
vllm_ascend/compilation/passes/sequence_parallelism_moe.py
Normal file
204
vllm_ascend/compilation/passes/sequence_parallelism_moe.py
Normal file
@@ -0,0 +1,204 @@
|
||||
import torch
|
||||
import torch._inductor.pattern_matcher as pm
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
from vllm.compilation.passes.vllm_inductor_pass import PatternPrettyPrinter, VllmInductorPass
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.utils import Range
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.compilation.passes.sequence_parallelism import (
|
||||
_SequenceParallelPatternHelper,
|
||||
get_sp_min_token_num,
|
||||
)
|
||||
|
||||
|
||||
class MiddleLayerAllgatherAddRMSNormPattern(_SequenceParallelPatternHelper):
|
||||
"""Replaces all_gather + slice + AddRMSNormBias with AddRMSNormBias +
|
||||
all_gather to avoid middle-layer shape mismatch."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def get_inputs(self):
|
||||
input = self.empty(5, 16)
|
||||
weight = self.empty(16)
|
||||
residual = self.empty(8, 16)
|
||||
# num_tokens = 8
|
||||
return [input, weight, residual]
|
||||
|
||||
def get_scalar_inputs(self):
|
||||
return {"num_tokens": 8}
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(
|
||||
input: torch.Tensor, weight: torch.Tensor, residual: torch.Tensor, num_tokens
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
all_gather = self._all_gather(input)
|
||||
x_sliced = all_gather[:num_tokens]
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(x_sliced, residual, weight, None, self.eps)
|
||||
|
||||
return result, residual
|
||||
|
||||
def replacement(
|
||||
input: torch.Tensor, weight: torch.Tensor, residual: torch.Tensor, num_tokens
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
residual = torch.ops.vllm.maybe_chunk_residual(input, residual)
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(input, residual, weight, None, self.eps)
|
||||
all_gather = self._all_gather(result)
|
||||
return all_gather, residual
|
||||
|
||||
pm.register_replacement(
|
||||
pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass, scalar_workaround=self.get_scalar_inputs()
|
||||
)
|
||||
|
||||
|
||||
class LastLayerAllgatherRMSNormPattern(_SequenceParallelPatternHelper):
|
||||
"""Same as MiddleLayerAllgatherAddRMSNormPattern but for the last layer (no residual)
|
||||
all_gather + RMSNorm fusion."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def get_inputs(self):
|
||||
input = self.empty(5, 16)
|
||||
weight = self.empty(16)
|
||||
residual = self.empty(8, 16)
|
||||
return [input, weight, residual]
|
||||
|
||||
def get_scalar_inputs(self):
|
||||
return {"num_tokens": 8}
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(
|
||||
input: torch.Tensor, weight: torch.Tensor, residual: torch.Tensor, num_tokens
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
all_gather = self._all_gather(input)
|
||||
x_sliced = all_gather[:num_tokens]
|
||||
result, _, _ = torch.ops._C_ascend.npu_add_rms_norm_bias(x_sliced, residual, weight, None, self.eps)
|
||||
|
||||
return result
|
||||
|
||||
def replacement(
|
||||
input: torch.Tensor, weight: torch.Tensor, residual: torch.Tensor, num_tokens
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
residual = torch.ops.vllm.maybe_chunk_residual(input, residual)
|
||||
result, _, _ = torch.ops._C_ascend.npu_add_rms_norm_bias(input, residual, weight, None, self.eps)
|
||||
all_gather = self._all_gather(result)
|
||||
return all_gather
|
||||
|
||||
pm.register_replacement(
|
||||
pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass, scalar_workaround=self.get_scalar_inputs()
|
||||
)
|
||||
|
||||
|
||||
class Qwen3VLMiddleLayerAllgatherAddRMSNormPattern(_SequenceParallelPatternHelper):
|
||||
"""Replaces all_gather + slice + add + AddRMSNormBias with add(chunk) +
|
||||
AddRMSNormBias + all_gather for Qwen3-VL-style all_gather path."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def get_inputs(self):
|
||||
input = self.empty(5, 16)
|
||||
weight = self.empty(16)
|
||||
residual = self.empty(8, 16)
|
||||
deepstack_input_embeds = self.empty(8, 16)
|
||||
return [input, weight, residual, deepstack_input_embeds]
|
||||
|
||||
def get_scalar_inputs(self):
|
||||
return {"num_tokens": 8}
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
deepstack_input_embeds: torch.Tensor,
|
||||
num_tokens,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
all_gather = self._all_gather(input)
|
||||
x_sliced = all_gather[:num_tokens]
|
||||
add_ = x_sliced + deepstack_input_embeds
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(add_, residual, weight, None, self.eps)
|
||||
|
||||
return result, residual
|
||||
|
||||
def replacement(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
deepstack_input_embeds: torch.Tensor,
|
||||
num_tokens,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
chunk = deepstack_input_embeds.chunk(self.tp_size)[self.tp_rank]
|
||||
add_ = input + chunk
|
||||
residual = torch.ops.vllm.maybe_chunk_residual(input, residual)
|
||||
result, _, residual = torch.ops._C_ascend.npu_add_rms_norm_bias(add_, residual, weight, None, self.eps)
|
||||
all_gather = self._all_gather(result)
|
||||
return all_gather, residual
|
||||
|
||||
pm.register_replacement(
|
||||
pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass, scalar_workaround=self.get_scalar_inputs()
|
||||
)
|
||||
|
||||
|
||||
class AllGatherChunkNoOpPattern(_SequenceParallelPatternHelper):
|
||||
"""Folds all_gather + sequence_parallel_chunk_impl into identity (no-op)."""
|
||||
|
||||
def __init__(self, vllm_config: VllmConfig, eps: float = 1e-6):
|
||||
super().__init__(eps, vllm_config.model_config.dtype, torch.npu.current_device())
|
||||
|
||||
def get_inputs(self):
|
||||
return [self.empty(8, 16)]
|
||||
|
||||
def register(self, pm_pass: PatternMatcherPass):
|
||||
def pattern(input: torch.Tensor) -> torch.Tensor:
|
||||
gathered = self._all_gather(input)
|
||||
return torch.ops.vllm.sequence_parallel_chunk_impl(gathered)
|
||||
|
||||
def replacement(input: torch.Tensor) -> torch.Tensor:
|
||||
return input
|
||||
|
||||
pm.register_replacement(pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass)
|
||||
|
||||
|
||||
class SequenceParallelismMoePass(VllmInductorPass):
|
||||
"""Sequence parallelism AllGather epilogue pass.
|
||||
|
||||
Applies AllGather-based patterns: MiddleLayerAllgatherAddRMSNormPattern,
|
||||
LastLayerAllgatherRMSNormPattern, Qwen3VLMiddleLayerAllgatherAddRMSNormPattern,
|
||||
and AllGatherChunkNoOpPattern (all_gather + sequence_parallel_chunk_impl -> identity).
|
||||
"""
|
||||
|
||||
def __init__(self, config: VllmConfig):
|
||||
super().__init__(config)
|
||||
|
||||
self.patterns: PatternMatcherPass = PatternMatcherPass(pass_name="npu_sequence_parallelism_allgather_ep_pass")
|
||||
|
||||
for epsilon in [1e-5, 1e-6]:
|
||||
MiddleLayerAllgatherAddRMSNormPattern(config, epsilon).register(self.patterns)
|
||||
LastLayerAllgatherRMSNormPattern(config, epsilon).register(self.patterns)
|
||||
Qwen3VLMiddleLayerAllgatherAddRMSNormPattern(config, epsilon).register(self.patterns)
|
||||
|
||||
AllGatherChunkNoOpPattern(config).register(self.patterns)
|
||||
|
||||
self.min_tokens = get_sp_min_token_num(config)
|
||||
|
||||
def __call__(self, graph: torch.fx.Graph):
|
||||
self.begin()
|
||||
logger.debug("before apply replacement %s", str(graph))
|
||||
self.matched_count = self.patterns.apply(graph)
|
||||
logger.debug("after apply replacement %s", str(graph))
|
||||
logger.debug("SequenceParallelismMoePass replaced %s patterns", self.matched_count)
|
||||
pattern_idx = 0
|
||||
for pattern_entry in self.patterns.patterns.values():
|
||||
for p in pattern_entry:
|
||||
p_str = PatternPrettyPrinter.run(p.pattern)
|
||||
logger.debug("Pattern %d: %s", pattern_idx, p_str)
|
||||
pattern_idx += 1
|
||||
self.end_and_log()
|
||||
|
||||
def is_applicable_for_range(self, compile_range: Range) -> bool:
|
||||
applicable = compile_range.start >= self.min_tokens
|
||||
logger.debug("SequenceParallelismMoePass compile_range=%r applicable=%r", compile_range, applicable)
|
||||
return applicable
|
||||
0
vllm_ascend/compilation/passes/utils/__init__.py
Normal file
0
vllm_ascend/compilation/passes/utils/__init__.py
Normal file
@@ -0,0 +1,75 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
|
||||
from torch._inductor.pattern_matcher import Match
|
||||
from vllm.logger import logger
|
||||
|
||||
|
||||
def extra_stream_scope_check(match: Match) -> bool:
|
||||
"""
|
||||
Checks if all nodes in the same stream.
|
||||
"""
|
||||
non_default_streams = set()
|
||||
has_default = False
|
||||
|
||||
for node in match.nodes:
|
||||
if node.op == "call_function":
|
||||
current_stream = node.meta.get("stream_label")
|
||||
if current_stream is None:
|
||||
has_default = True
|
||||
else:
|
||||
non_default_streams.add(current_stream)
|
||||
if len(non_default_streams) > 1:
|
||||
logger.debug(
|
||||
"Cross-stream operation detected in pattern match for AddRMSNormQuant. "
|
||||
"Multiple streams found: %s. Fusion is not supported for cross-stream operations.",
|
||||
non_default_streams,
|
||||
)
|
||||
return False
|
||||
|
||||
if has_default and len(non_default_streams) > 0:
|
||||
logger.debug(
|
||||
"Cross-stream operation detected in pattern match for AddRMSNormQuant. "
|
||||
"Multiple streams found: %s. Fusion is not supported for cross-stream operations.",
|
||||
non_default_streams,
|
||||
)
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
_register_patterns = set()
|
||||
|
||||
|
||||
def check_and_register_fusion_pass(pattern_class: type, **kwargs):
|
||||
global _register_patterns
|
||||
eps = kwargs.get("eps", 1e-6)
|
||||
pattern_key = str(pattern_class.__name__) + str(eps)
|
||||
if pattern_key in _register_patterns:
|
||||
return
|
||||
|
||||
pattern = pattern_class(**kwargs)
|
||||
try:
|
||||
pattern.register()
|
||||
_register_patterns.add(pattern_key)
|
||||
except RuntimeError as e:
|
||||
if "Duplicate pattern" in str(e):
|
||||
logger.warning("Pattern %s eps %s has been registered", pattern_class.__name__, eps)
|
||||
_register_patterns.add(pattern_key)
|
||||
else:
|
||||
raise e
|
||||
Reference in New Issue
Block a user