369 lines
16 KiB
Python
369 lines
16 KiB
Python
#
|
|
# 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 copy
|
|
import functools
|
|
import os
|
|
from collections.abc import Callable
|
|
from typing import Any, cast
|
|
|
|
import torch
|
|
import torch.fx as fx
|
|
from torch._dynamo.backends.common import aot_autograd
|
|
from torch._inductor.compile_fx import graph_returns_tuple, make_graph_return_tuple
|
|
from torch._inductor.decomposition import select_decomp_table
|
|
from torch.fx import GraphModule
|
|
from vllm.compilation.compiler_interface import CompilerInterface
|
|
from vllm.config import VllmConfig
|
|
from vllm.config.utils import Range
|
|
from vllm.logger import logger
|
|
|
|
from vllm_ascend.ascend_config import AscendCompilationConfig, get_ascend_config
|
|
from vllm_ascend.utils import COMPILATION_PASS_KEY
|
|
|
|
|
|
def compile_fx(graph: GraphModule, example_inputs: list, inner_compile: Callable, decompositions: dict) -> Callable:
|
|
recursive_compile_fx = functools.partial(compile_fx, inner_compile=inner_compile, decompositions=decompositions)
|
|
|
|
if not graph_returns_tuple(graph):
|
|
return make_graph_return_tuple(graph, example_inputs, recursive_compile_fx)
|
|
return aot_autograd(fw_compiler=inner_compile)(graph, example_inputs)
|
|
|
|
|
|
def fusion_pass_compile(
|
|
graph: fx.GraphModule,
|
|
example_inputs: list[Any],
|
|
compiler_config: dict[str, Any],
|
|
compile_range: Range,
|
|
key: str | None = None,
|
|
) -> tuple[Callable | None, Any | None]:
|
|
def compile_inner(graph, example_inputs):
|
|
current_pass_manager = compiler_config[COMPILATION_PASS_KEY]
|
|
graph = current_pass_manager(graph)
|
|
return graph
|
|
|
|
decompositions = select_decomp_table()
|
|
|
|
compiled_fn = compile_fx(
|
|
graph=graph,
|
|
example_inputs=example_inputs,
|
|
inner_compile=compile_inner,
|
|
decompositions=decompositions,
|
|
)
|
|
|
|
return compiled_fn, None
|
|
|
|
|
|
def _compute_decode_cudagraph_batch_sizes(vllm_config: VllmConfig) -> list[int]:
|
|
num_spec_tokens = vllm_config.speculative_config.num_speculative_tokens if vllm_config.speculative_config else 0
|
|
uniform_decode_query_len = num_spec_tokens + 1
|
|
max_num_tokens = vllm_config.scheduler_config.max_num_seqs * uniform_decode_query_len
|
|
return [
|
|
x
|
|
for x in vllm_config.compilation_config.cudagraph_capture_sizes
|
|
if max_num_tokens >= x >= uniform_decode_query_len
|
|
]
|
|
|
|
|
|
def _configure_backend(
|
|
config: Any,
|
|
ascend_compilation_config: AscendCompilationConfig,
|
|
vllm_config: VllmConfig,
|
|
process_kwargs_options: Callable | None = None,
|
|
) -> None:
|
|
if ascend_compilation_config.enable_static_kernel:
|
|
# npugraph_ex's static_kernel requires LOCAL_WORLD_SIZE to determine the
|
|
# physical node topology for creating per-node Gloo groups, which
|
|
# coordinate static kernel compilation and .run package installation.
|
|
# vLLM does not set this env var by default (unlike torchrun), so we
|
|
# compute it from parallel config:
|
|
# local_world_size: processes per node for one DP replica
|
|
# data_parallel_size_local: number of DP replicas on this node
|
|
# actual_local_world_size: total processes on this physical machine
|
|
if "LOCAL_WORLD_SIZE" not in os.environ:
|
|
actual_local_world_size = (
|
|
vllm_config.parallel_config.local_world_size * vllm_config.parallel_config.data_parallel_size_local
|
|
)
|
|
os.environ["LOCAL_WORLD_SIZE"] = str(actual_local_world_size)
|
|
logger.info_once(
|
|
"Setting LOCAL_WORLD_SIZE=%d for static kernel (local_world_size=%d * data_parallel_size_local=%d).",
|
|
actual_local_world_size,
|
|
vllm_config.parallel_config.local_world_size,
|
|
vllm_config.parallel_config.data_parallel_size_local,
|
|
scope="global",
|
|
)
|
|
|
|
if process_kwargs_options is not None:
|
|
# npugraph_ex (both old and new): build options dict and use _process_kwargs_options.
|
|
# It maps flat option names to nested config paths for old versions,
|
|
# and directly setattr for new versions with flat CompilerConfig.
|
|
# force_eager=True: execute FX graph in eager mode before graph capture.
|
|
# inplace_pass=False: disable reinplace pass to avoid gelu fallback to CPU.
|
|
options: dict[str, Any] = {
|
|
"force_eager": True,
|
|
"inplace_pass": False,
|
|
"clone_input": False,
|
|
"clone_output": False,
|
|
}
|
|
if ascend_compilation_config.enable_static_kernel:
|
|
logger.info_once(
|
|
"enable_static_kernel is enabled, static shape kernel will be used to accelerate aclgraph execution.",
|
|
scope="global",
|
|
)
|
|
options["static_kernel_compile"] = True
|
|
# Set sym_range to limit static kernel compilation to specified batch sizes.
|
|
options["_vllm_aclnn_static_kernel_sym_range"] = _compute_decode_cudagraph_batch_sizes(vllm_config)
|
|
process_kwargs_options(config, {"options": options})
|
|
else:
|
|
# torchair (reduce-overhead): use nested config structure directly.
|
|
# mode="reduce-overhead": use aclgraph mode, avoid fx graph to Ascend IR transformation.
|
|
config.mode = "reduce-overhead"
|
|
config.debug.run_eagerly = True
|
|
# Disable reinplace pass to avoid gelu fallback to CPU causing host-device copy error.
|
|
config.debug.aclgraph.disable_reinplace_inplaceable_ops_pass = True
|
|
if ascend_compilation_config.enable_static_kernel:
|
|
logger.info_once(
|
|
"enable_static_kernel is enabled, static shape kernel will be used to accelerate aclgraph execution.",
|
|
scope="global",
|
|
)
|
|
config.experimental_config.aclgraph._aclnn_static_shape_kernel = True
|
|
config.experimental_config.aclgraph._aclnn_static_shape_kernel_sym_value_range = (
|
|
_compute_decode_cudagraph_batch_sizes(vllm_config)
|
|
)
|
|
|
|
|
|
def npugraph_ex_compile(
|
|
graph: fx.GraphModule,
|
|
example_inputs: list[Any],
|
|
compiler_config: dict[str, Any],
|
|
vllm_config: VllmConfig,
|
|
ascend_compilation_config: AscendCompilationConfig,
|
|
compile_range: Range,
|
|
key: str | None = None,
|
|
cache_dir: str | None = None,
|
|
) -> tuple[Callable | None, Any | None]:
|
|
# Try npugraph_ex first, fall back to torchair for backward compatibility.
|
|
try:
|
|
import npugraph_ex as nge
|
|
|
|
cache_path = os.path.join(cache_dir, key) if (cache_dir and key) else None
|
|
torch.npu.set_compile_mode(jit_compile=False)
|
|
config = nge.CompilerConfig()
|
|
# _process_kwargs_options exists in both old and new npugraph_ex,
|
|
# but in different modules: new -> compiler_config, old -> npugraphex_config.
|
|
try:
|
|
from npugraph_ex.configs.compiler_config import _process_kwargs_options
|
|
except ImportError:
|
|
from npugraph_ex.configs.npugraphex_config import _process_kwargs_options
|
|
_configure_backend(
|
|
config, ascend_compilation_config, vllm_config, process_kwargs_options=_process_kwargs_options
|
|
)
|
|
import npugraph_ex.npu_fx_compiler as nfx
|
|
|
|
_original_get_compiled_gm = nfx._NpuFxCompiler._get_compiled_gm
|
|
|
|
def patched_get_compiled_gm(self, graph, example_inputs):
|
|
compiled_gm = _original_get_compiled_gm(self, graph, example_inputs)
|
|
if cache_path:
|
|
py_code = compiled_gm.get_code()
|
|
if py_code:
|
|
# Triton kernel indices (kernel_side_table) are registered in-process
|
|
# at compile time and are not serializable across process boundaries.
|
|
# Graphs containing triton_kernel_wrapper calls must not be cached,
|
|
# because loading the py_code in a new process will hit an
|
|
# AssertionError in kernel_side_table.get_kernel().
|
|
if "triton_kernel_wrapper" in py_code:
|
|
logger.info(
|
|
"Skipping npugraph_ex cache for graph containing Triton kernels "
|
|
"(kernel_side_table indices are process-local): %s",
|
|
cache_path,
|
|
)
|
|
else:
|
|
os.makedirs(os.path.dirname(cache_path), exist_ok=True)
|
|
with open(cache_path, "w") as f:
|
|
f.write(py_code)
|
|
logger.info("Saved compiled graph to cache: %s", cache_path)
|
|
return compiled_gm
|
|
|
|
nfx._NpuFxCompiler._get_compiled_gm = patched_get_compiled_gm
|
|
backend = nge.get_npu_backend(compiler_config=config)
|
|
# torch.compile requires the output of the fx graph to be a tuple
|
|
if not graph_returns_tuple(graph):
|
|
compiled_fn = make_graph_return_tuple(graph, example_inputs, backend)
|
|
else:
|
|
compiled_fn = backend(graph, example_inputs)
|
|
nfx._NpuFxCompiler._get_compiled_gm = _original_get_compiled_gm
|
|
return compiled_fn, (key, cache_path)
|
|
except ImportError:
|
|
import torchair
|
|
|
|
torch.npu.set_compile_mode(jit_compile=False)
|
|
config = torchair.CompilerConfig()
|
|
_configure_backend(config, ascend_compilation_config, vllm_config)
|
|
backend = torchair.get_npu_backend(compiler_config=config)
|
|
# torch.compile requires the output of the fx graph to be a tuple
|
|
if not graph_returns_tuple(graph):
|
|
compiled_fn = make_graph_return_tuple(graph, example_inputs, backend)
|
|
else:
|
|
compiled_fn = backend(graph, example_inputs)
|
|
return compiled_fn, None
|
|
|
|
|
|
class AscendCompiler(CompilerInterface):
|
|
"""
|
|
AscendCompiler is a custom compiler interface for the Ascend platform.
|
|
This class provides a method to compile a PyTorch FX graph module with
|
|
specific configurations for graph fusion and decomposition.
|
|
"""
|
|
|
|
name = "AscendCompiler"
|
|
|
|
# TODO(wxs): add passes related to compilation in compute_hash
|
|
def compute_hash(self, vllm_config: VllmConfig) -> str:
|
|
self.vllm_config = vllm_config
|
|
ascend_compilation_config = get_ascend_config().ascend_compilation_config
|
|
from hashlib import sha256
|
|
|
|
import torch_npu
|
|
|
|
factors = {
|
|
"torch_npu_version": torch_npu.__version__,
|
|
"enable_npugraph_ex": ascend_compilation_config.enable_npugraph_ex,
|
|
"enable_static_kernel": ascend_compilation_config.enable_static_kernel,
|
|
}
|
|
logger.info("AscendCompiler hash factors: %s", factors)
|
|
return sha256(str(factors).encode(), usedforsecurity=False).hexdigest()[:10]
|
|
|
|
def initialize_cache(self, cache_dir, disable_cache=False, prefix=""):
|
|
self.cache_dir = cache_dir
|
|
self.disable_cache = disable_cache
|
|
|
|
def compile(
|
|
self,
|
|
graph: fx.GraphModule,
|
|
example_inputs: list[Any],
|
|
compiler_config: dict[str, Any],
|
|
compile_range: Range,
|
|
key: str | None = None,
|
|
) -> tuple[Callable | None, Any | None]:
|
|
# inductor can inplace modify the graph, so we need to copy it
|
|
# see https://github.com/pytorch/pytorch/issues/138980
|
|
graph = copy.deepcopy(graph)
|
|
|
|
from torch._guards import detect_fake_mode
|
|
|
|
current_fake_mode = detect_fake_mode()
|
|
if current_fake_mode is not None:
|
|
example_inputs = [
|
|
current_fake_mode.from_tensor(inp)
|
|
if (
|
|
isinstance(inp, torch.Tensor)
|
|
and hasattr(inp, "fake_mode")
|
|
and inp.fake_mode is not current_fake_mode
|
|
)
|
|
else inp
|
|
for inp in example_inputs
|
|
]
|
|
|
|
ascend_compilation_config = get_ascend_config().ascend_compilation_config
|
|
if ascend_compilation_config.enable_npugraph_ex:
|
|
cache_dir = None if getattr(self, "disable_cache", False) else getattr(self, "cache_dir", None)
|
|
logger.info_once(
|
|
"enable_npugraph_ex is enabled, which will bring graph compilation optimization.",
|
|
scope="global",
|
|
)
|
|
assert hasattr(self, "vllm_config")
|
|
return npugraph_ex_compile(
|
|
graph,
|
|
example_inputs,
|
|
compiler_config,
|
|
self.vllm_config,
|
|
ascend_compilation_config,
|
|
compile_range,
|
|
key,
|
|
cache_dir,
|
|
)
|
|
else:
|
|
return fusion_pass_compile(graph, example_inputs, compiler_config, compile_range, key)
|
|
|
|
def load(self, handle, graph, example_inputs, graph_index, compile_range):
|
|
key, path = handle
|
|
# Cache file may be absent when the graph was skipped at save time (e.g. it
|
|
# contained Triton kernels whose kernel_side_table indices are process-local
|
|
# and cannot be serialized). Fall back to a fresh compilation so the Triton
|
|
# kernels are properly registered in the current process.
|
|
if not path or not os.path.exists(path):
|
|
logger.info(
|
|
"npugraph_ex cache miss for key %s (file absent or not saved), recompiling",
|
|
key,
|
|
)
|
|
# Mirror the same pre-processing done in compile(): deepcopy the graph
|
|
# to prevent make_graph_return_tuple from mutating the caller's copy,
|
|
# and re-wrap FakeTensor inputs under the current fake mode to avoid
|
|
# "fake mode mismatch" in aot_module_simplified.
|
|
graph = copy.deepcopy(graph)
|
|
from torch._guards import detect_fake_mode
|
|
|
|
current_fake_mode = detect_fake_mode()
|
|
if current_fake_mode is not None:
|
|
example_inputs = [
|
|
current_fake_mode.from_tensor(inp)
|
|
if (
|
|
isinstance(inp, torch.Tensor)
|
|
and hasattr(inp, "fake_mode")
|
|
and inp.fake_mode is not current_fake_mode
|
|
)
|
|
else inp
|
|
for inp in example_inputs
|
|
]
|
|
ascend_compilation_config = get_ascend_config().ascend_compilation_config
|
|
assert hasattr(self, "vllm_config")
|
|
compiled_fn, _ = npugraph_ex_compile(
|
|
graph,
|
|
example_inputs,
|
|
{},
|
|
self.vllm_config,
|
|
ascend_compilation_config,
|
|
compile_range,
|
|
key,
|
|
getattr(self, "cache_dir", None),
|
|
)
|
|
return compiled_fn
|
|
|
|
from npugraph_ex.npu_fx_compiler import _CompiledFxArtifacts, _CompiledFxGraph
|
|
|
|
with open(path) as f:
|
|
py_code = f.read()
|
|
artifacts = _CompiledFxArtifacts()
|
|
artifacts.py_code = py_code
|
|
logger.info("Loaded npugraph_ex compilation cache from %s", path)
|
|
compiled_fn = cast(Callable[..., Any], _CompiledFxGraph.load_artifacts(artifacts))
|
|
|
|
# The saved code was compiled from the graph after make_graph_return_tuple mutated it
|
|
# to return a flat tuple. If the original graph didn't return a tuple, we need to
|
|
# recreate the unflatten wrapper so callers receive the original output structure.
|
|
if not graph_returns_tuple(graph):
|
|
_inner_fn = compiled_fn
|
|
|
|
def compiled_fn(*args, **kwargs):
|
|
result = _inner_fn(*args, **kwargs)
|
|
if isinstance(result, (tuple, list)) and len(result) == 1:
|
|
return result[0]
|
|
return result
|
|
|
|
return compiled_fn
|