forked from Karylab-cklius/vllm
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8f4f8425d0 |
@@ -247,3 +247,92 @@ def test_model_startup(monkeypatch, vllm_runner, fresh_vllm_cache, spec):
|
|||||||
|
|
||||||
# Warm start — compiled artifacts loaded from disk cache.
|
# Warm start — compiled artifacts loaded from disk cache.
|
||||||
_check_model_run(vllm_runner, spec, is_cold_start=False)
|
_check_model_run(vllm_runner, spec, is_cold_start=False)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# compile_model (compile-only) cold start tests
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
COMPILE_ONLY_SPECS = [
|
||||||
|
pytest.param(
|
||||||
|
ModelStartupSpec(
|
||||||
|
model="microsoft/Phi-tiny-MoE-instruct",
|
||||||
|
hf_overrides={},
|
||||||
|
cold_artifacts_saved=3,
|
||||||
|
warm_artifacts_saved=0,
|
||||||
|
warm_artifacts_loaded=3,
|
||||||
|
),
|
||||||
|
id="phi_tiny_moe",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
ModelStartupSpec(
|
||||||
|
model="openai/gpt-oss-120b",
|
||||||
|
hf_overrides={
|
||||||
|
"num_hidden_layers": 8,
|
||||||
|
"hidden_size": 256,
|
||||||
|
"intermediate_size": 512,
|
||||||
|
"num_attention_heads": 8,
|
||||||
|
"num_key_value_heads": 1,
|
||||||
|
"num_local_experts": 8,
|
||||||
|
},
|
||||||
|
cold_artifacts_saved=3,
|
||||||
|
warm_artifacts_saved=0,
|
||||||
|
warm_artifacts_loaded=3,
|
||||||
|
),
|
||||||
|
id="gpt_oss_120b",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
ModelStartupSpec(
|
||||||
|
model="zai-org/GLM-4.5",
|
||||||
|
hf_overrides=_SMALL_MOE_OVERRIDES,
|
||||||
|
cold_artifacts_saved=4,
|
||||||
|
warm_artifacts_saved=0,
|
||||||
|
warm_artifacts_loaded=4,
|
||||||
|
),
|
||||||
|
id="glm_4.5",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _compile_only_cold_start(spec: ModelStartupSpec):
|
||||||
|
"""Cold start using compile_model (fake weights, no GPU memory)."""
|
||||||
|
from vllm.compile_only import compile_model
|
||||||
|
|
||||||
|
old = compilation_counter.clone()
|
||||||
|
compile_model(
|
||||||
|
spec.model,
|
||||||
|
trust_remote_code=True,
|
||||||
|
max_model_len=256,
|
||||||
|
max_num_batched_tokens=1024,
|
||||||
|
block_size=64,
|
||||||
|
hf_overrides=spec.hf_overrides,
|
||||||
|
compilation_config=CompilationConfig(
|
||||||
|
mode=CompilationMode.VLLM_COMPILE,
|
||||||
|
cudagraph_mode=CUDAGraphMode.NONE,
|
||||||
|
pass_config=PassConfig(fuse_allreduce_rms=False),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
saved = (
|
||||||
|
compilation_counter.num_compiled_artifacts_saved
|
||||||
|
- old.num_compiled_artifacts_saved
|
||||||
|
)
|
||||||
|
print(f"\n=== COMPILE-ONLY COLD START for {spec.model} ===")
|
||||||
|
print(f" num_compiled_artifacts_saved={saved}")
|
||||||
|
assert saved == spec.cold_artifacts_saved, f"cold_artifacts_saved: got {saved}"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("spec", COMPILE_ONLY_SPECS)
|
||||||
|
@fork_new_process_for_each_test
|
||||||
|
def test_compile_only_startup(monkeypatch, vllm_runner, fresh_vllm_cache, spec):
|
||||||
|
monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
|
||||||
|
|
||||||
|
# Cold start: compile-only in a forked child (fork before CUDA init).
|
||||||
|
ctx = mp.get_context("fork")
|
||||||
|
p = ctx.Process(target=_compile_only_cold_start, args=(spec,))
|
||||||
|
p.start()
|
||||||
|
p.join()
|
||||||
|
assert p.exitcode == 0, "Compile-only cold start failed"
|
||||||
|
|
||||||
|
# Warm start — compiled artifacts loaded from disk cache.
|
||||||
|
_check_model_run(vllm_runner, spec, is_cold_start=False)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import typing
|
|||||||
import vllm.env_override # noqa: F401
|
import vllm.env_override # noqa: F401
|
||||||
|
|
||||||
MODULE_ATTRS = {
|
MODULE_ATTRS = {
|
||||||
|
"compile_model": ".compile_only:compile_model",
|
||||||
"AsyncEngineArgs": ".engine.arg_utils:AsyncEngineArgs",
|
"AsyncEngineArgs": ".engine.arg_utils:AsyncEngineArgs",
|
||||||
"EngineArgs": ".engine.arg_utils:EngineArgs",
|
"EngineArgs": ".engine.arg_utils:EngineArgs",
|
||||||
"AsyncLLMEngine": ".engine.async_llm_engine:AsyncLLMEngine",
|
"AsyncLLMEngine": ".engine.async_llm_engine:AsyncLLMEngine",
|
||||||
@@ -39,6 +40,7 @@ MODULE_ATTRS = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if typing.TYPE_CHECKING:
|
if typing.TYPE_CHECKING:
|
||||||
|
from vllm.compile_only import compile_model as compile_model
|
||||||
from vllm.engine.arg_utils import AsyncEngineArgs, EngineArgs
|
from vllm.engine.arg_utils import AsyncEngineArgs, EngineArgs
|
||||||
from vllm.engine.async_llm_engine import AsyncLLMEngine
|
from vllm.engine.async_llm_engine import AsyncLLMEngine
|
||||||
from vllm.engine.llm_engine import LLMEngine
|
from vllm.engine.llm_engine import LLMEngine
|
||||||
|
|||||||
@@ -50,6 +50,16 @@ from .passes.pass_manager import PostGradPassManager
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class CompilationDone(Exception):
|
||||||
|
"""Raised in compile-only mode after compilation is complete.
|
||||||
|
|
||||||
|
This signals that the vLLM-compile cache has been populated and
|
||||||
|
there is no need to actually execute the compiled code.
|
||||||
|
"""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def make_copy_and_call(
|
def make_copy_and_call(
|
||||||
sym_tensor_indices: list[int],
|
sym_tensor_indices: list[int],
|
||||||
input_buffers: list[torch.Tensor | None],
|
input_buffers: list[torch.Tensor | None],
|
||||||
@@ -990,7 +1000,14 @@ class VllmBackend:
|
|||||||
# Compute config/compiler/code hashes once and reuse
|
# Compute config/compiler/code hashes once and reuse
|
||||||
config_hash = vllm_config.compute_hash()
|
config_hash = vllm_config.compute_hash()
|
||||||
compiler_hash = self.compiler_manager.compute_hash(vllm_config)
|
compiler_hash = self.compiler_manager.compute_hash(vllm_config)
|
||||||
forward_code_files = list(sorted(self.compilation_config.traced_files))
|
# Filter out PyTorch internal files — they are already covered
|
||||||
|
# by the torch version in env_factors.
|
||||||
|
torch_root = os.path.dirname(torch.__file__) + os.sep
|
||||||
|
forward_code_files = [
|
||||||
|
f
|
||||||
|
for f in sorted(self.compilation_config.traced_files)
|
||||||
|
if not f.startswith(torch_root)
|
||||||
|
]
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Traced files (to be considered for compilation cache):\n%s",
|
"Traced files (to be considered for compilation cache):\n%s",
|
||||||
|
|||||||
@@ -599,9 +599,23 @@ def _support_torch_compile(
|
|||||||
# AOT artifact.
|
# AOT artifact.
|
||||||
self.save_aot_compiled_function()
|
self.save_aot_compiled_function()
|
||||||
|
|
||||||
|
# In compile-only mode, raise CompilationDone after all
|
||||||
|
# piecewise graphs are compiled and cache artifacts saved.
|
||||||
|
# This is caught in gpu_worker.compile_or_warm_up_model()
|
||||||
|
# to skip execution with fake tensors.
|
||||||
|
if self.compilation_config.compile_only:
|
||||||
|
from .backends import CompilationDone
|
||||||
|
|
||||||
|
raise CompilationDone
|
||||||
|
|
||||||
with monitor_profiling_run():
|
with monitor_profiling_run():
|
||||||
output = self.aot_compiled_fn(self, *args, **kwargs)
|
output = self.aot_compiled_fn(self, *args, **kwargs)
|
||||||
else:
|
else:
|
||||||
|
# Same as above for non-AOT path.
|
||||||
|
if self.compilation_config.compile_only:
|
||||||
|
from .backends import CompilationDone
|
||||||
|
|
||||||
|
raise CompilationDone
|
||||||
with monitor_torch_compile(
|
with monitor_torch_compile(
|
||||||
self.vllm_config,
|
self.vllm_config,
|
||||||
"torch.compile and initial profiling/warmup "
|
"torch.compile and initial profiling/warmup "
|
||||||
|
|||||||
@@ -0,0 +1,90 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||||
|
"""Compile-only mode: populate vLLM's torch.compile cache
|
||||||
|
without loading real model weights or allocating KV caches.
|
||||||
|
|
||||||
|
The compile-only flag causes the model loader to be wrapped with
|
||||||
|
FakeTensorMode (see ``fake_loader.wrap_loader_with_fake``), so the
|
||||||
|
user's original ``load_format`` is preserved and the real loader's
|
||||||
|
full pipeline runs — just with fake tensors instead of real weights.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
from vllm.logger import init_logger
|
||||||
|
from vllm.usage.usage_lib import UsageContext
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def compile_model(
|
||||||
|
model: str,
|
||||||
|
*,
|
||||||
|
tensor_parallel_size: int = 1,
|
||||||
|
pipeline_parallel_size: int = 1,
|
||||||
|
quantization: str | None = None,
|
||||||
|
dtype: str = "auto",
|
||||||
|
trust_remote_code: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
"""Pre-populate vLLM's torch.compile cache for a model.
|
||||||
|
|
||||||
|
Runs compilation using fake weights (zero GPU memory)
|
||||||
|
so that vLLM's torch.compile cache is populated. Subsequent
|
||||||
|
``vllm serve`` or ``LLM(...)`` calls for the same model
|
||||||
|
configuration will hit the warm cache and skip compilation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: HuggingFace model name or path.
|
||||||
|
tensor_parallel_size: Number of tensor parallel GPUs.
|
||||||
|
pipeline_parallel_size: Number of pipeline parallel stages.
|
||||||
|
quantization: Quantization method (e.g. "fp8").
|
||||||
|
dtype: Model dtype.
|
||||||
|
trust_remote_code: Trust remote code from HuggingFace.
|
||||||
|
**kwargs: Additional arguments passed to ``EngineArgs``.
|
||||||
|
"""
|
||||||
|
from vllm.engine.arg_utils import EngineArgs
|
||||||
|
|
||||||
|
engine_args = EngineArgs(
|
||||||
|
model=model,
|
||||||
|
tensor_parallel_size=tensor_parallel_size,
|
||||||
|
pipeline_parallel_size=pipeline_parallel_size,
|
||||||
|
quantization=quantization,
|
||||||
|
dtype=dtype,
|
||||||
|
trust_remote_code=trust_remote_code,
|
||||||
|
enforce_eager=False,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
vllm_config = engine_args.create_engine_config(usage_context=UsageContext.LLM_CLASS)
|
||||||
|
vllm_config.compilation_config.compile_only = True
|
||||||
|
|
||||||
|
_run_compile_with_config(vllm_config)
|
||||||
|
|
||||||
|
|
||||||
|
def run_compile_only(args: argparse.Namespace) -> None:
|
||||||
|
"""Run compile-only mode from CLI arguments."""
|
||||||
|
from vllm.engine.arg_utils import EngineArgs
|
||||||
|
|
||||||
|
engine_args = EngineArgs.from_cli_args(args)
|
||||||
|
engine_args.enforce_eager = False
|
||||||
|
|
||||||
|
vllm_config = engine_args.create_engine_config(usage_context=UsageContext.LLM_CLASS)
|
||||||
|
vllm_config.compilation_config.compile_only = True
|
||||||
|
|
||||||
|
_run_compile_with_config(vllm_config)
|
||||||
|
|
||||||
|
|
||||||
|
def _run_compile_with_config(vllm_config) -> None:
|
||||||
|
"""Shared compile-only logic."""
|
||||||
|
from vllm.plugins import load_general_plugins
|
||||||
|
from vllm.v1.executor import Executor
|
||||||
|
|
||||||
|
load_general_plugins()
|
||||||
|
|
||||||
|
executor_class = Executor.get_class(vllm_config)
|
||||||
|
executor = executor_class(vllm_config)
|
||||||
|
executor.collective_rpc("compile_or_warm_up_model")
|
||||||
|
|
||||||
|
logger.info("Compile-only mode complete. Cache populated.")
|
||||||
|
executor.shutdown()
|
||||||
@@ -655,6 +655,12 @@ class CompilationConfig:
|
|||||||
local_cache_dir: str = field(default=None, init=False) # type: ignore
|
local_cache_dir: str = field(default=None, init=False) # type: ignore
|
||||||
"""local cache dir for each rank"""
|
"""local cache dir for each rank"""
|
||||||
|
|
||||||
|
compile_only: bool = False
|
||||||
|
"""If True, run in compile-only mode: only torch.compile
|
||||||
|
compilation, skip CUDA graph capture, kernel warmup, and sampler
|
||||||
|
warmup. Used to pre-populate vLLM's torch.compile cache without
|
||||||
|
allocating KV caches or setting up the full engine."""
|
||||||
|
|
||||||
fast_moe_cold_start: bool | None = None
|
fast_moe_cold_start: bool | None = None
|
||||||
"""Optimization for fast MOE cold start.
|
"""Optimization for fast MOE cold start.
|
||||||
|
|
||||||
@@ -739,6 +745,8 @@ class CompilationConfig:
|
|||||||
"static_forward_context",
|
"static_forward_context",
|
||||||
"pass_config", # handled separately below
|
"pass_config", # handled separately below
|
||||||
"dynamic_shapes_config", # handled separately below
|
"dynamic_shapes_config", # handled separately below
|
||||||
|
# compile_only doesn't affect the compiled graph
|
||||||
|
"compile_only",
|
||||||
}
|
}
|
||||||
|
|
||||||
from vllm.config.utils import get_hash_factors, hash_factors
|
from vllm.config.utils import get_hash_factors, hash_factors
|
||||||
|
|||||||
@@ -0,0 +1,51 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
from vllm.entrypoints.cli.types import CLISubcommand
|
||||||
|
from vllm.entrypoints.openai.cli_args import make_arg_parser
|
||||||
|
from vllm.entrypoints.utils import VLLM_SUBCMD_PARSER_EPILOG
|
||||||
|
from vllm.utils.argparse_utils import FlexibleArgumentParser
|
||||||
|
|
||||||
|
DESCRIPTION = """[Experimental] Populate vLLM's torch.compile cache for a model.
|
||||||
|
|
||||||
|
This command is experimental and a work in progress. Not all models and
|
||||||
|
configurations are supported yet.
|
||||||
|
|
||||||
|
This runs compilation using fake weights (zero GPU memory) so that
|
||||||
|
vLLM's torch.compile cache is populated. Subsequent ``vllm serve`` or
|
||||||
|
``LLM(...)`` calls for the same model will hit the warm cache and skip
|
||||||
|
compilation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class CompileSubcommand(CLISubcommand):
|
||||||
|
"""The ``compile`` subcommand for the vLLM CLI."""
|
||||||
|
|
||||||
|
name = "compile"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def cmd(args: argparse.Namespace) -> None:
|
||||||
|
from vllm.compile_only import run_compile_only
|
||||||
|
|
||||||
|
if hasattr(args, "model_tag") and args.model_tag is not None:
|
||||||
|
args.model = args.model_tag
|
||||||
|
run_compile_only(args)
|
||||||
|
|
||||||
|
def subparser_init(
|
||||||
|
self, subparsers: argparse._SubParsersAction
|
||||||
|
) -> FlexibleArgumentParser:
|
||||||
|
compile_parser = subparsers.add_parser(
|
||||||
|
self.name,
|
||||||
|
help="[Experimental] Populate vLLM's torch.compile cache for a model.",
|
||||||
|
description=DESCRIPTION,
|
||||||
|
usage="vllm compile [model_tag] [options]",
|
||||||
|
)
|
||||||
|
compile_parser = make_arg_parser(compile_parser)
|
||||||
|
compile_parser.epilog = VLLM_SUBCMD_PARSER_EPILOG.format(subcmd=self.name)
|
||||||
|
return compile_parser
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_init() -> list[CLISubcommand]:
|
||||||
|
return [CompileSubcommand()]
|
||||||
@@ -16,6 +16,7 @@ logger = init_logger(__name__)
|
|||||||
def main():
|
def main():
|
||||||
import vllm.entrypoints.cli.benchmark.main
|
import vllm.entrypoints.cli.benchmark.main
|
||||||
import vllm.entrypoints.cli.collect_env
|
import vllm.entrypoints.cli.collect_env
|
||||||
|
import vllm.entrypoints.cli.compile
|
||||||
import vllm.entrypoints.cli.launch
|
import vllm.entrypoints.cli.launch
|
||||||
import vllm.entrypoints.cli.openai
|
import vllm.entrypoints.cli.openai
|
||||||
import vllm.entrypoints.cli.run_batch
|
import vllm.entrypoints.cli.run_batch
|
||||||
@@ -26,6 +27,7 @@ def main():
|
|||||||
CMD_MODULES = [
|
CMD_MODULES = [
|
||||||
vllm.entrypoints.cli.openai,
|
vllm.entrypoints.cli.openai,
|
||||||
vllm.entrypoints.cli.serve,
|
vllm.entrypoints.cli.serve,
|
||||||
|
vllm.entrypoints.cli.compile,
|
||||||
vllm.entrypoints.cli.launch,
|
vllm.entrypoints.cli.launch,
|
||||||
vllm.entrypoints.cli.benchmark.main,
|
vllm.entrypoints.cli.benchmark.main,
|
||||||
vllm.entrypoints.cli.collect_env,
|
vllm.entrypoints.cli.collect_env,
|
||||||
|
|||||||
@@ -0,0 +1,131 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||||
|
"""Fake weight loader for compile-only mode.
|
||||||
|
|
||||||
|
Initializes the model on the meta device (preserving Parameter subclasses
|
||||||
|
like ModelWeightParameter), and runs process_weights_after_loading on
|
||||||
|
meta tensors.
|
||||||
|
|
||||||
|
This is used for compile-only mode where we want to run torch.compile
|
||||||
|
without actually allocating any GPU memory for the model.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from vllm.config import ModelConfig, VllmConfig
|
||||||
|
from vllm.logger import init_logger
|
||||||
|
from vllm.model_executor.model_loader.base_loader import BaseModelLoader
|
||||||
|
from vllm.model_executor.model_loader.utils import (
|
||||||
|
initialize_model,
|
||||||
|
process_weights_after_loading,
|
||||||
|
)
|
||||||
|
from vllm.utils.torch_utils import set_default_torch_dtype
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeModelLoader(BaseModelLoader):
|
||||||
|
"""Model loader that initializes on meta device.
|
||||||
|
|
||||||
|
Model initialization runs on ``meta`` device because FakeTensorMode
|
||||||
|
doesn't preserve Parameter subclasses (e.g. ``ModelWeightParameter``
|
||||||
|
becomes a plain ``FakeTensor``). No GPU memory is allocated.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def download_model(self, model_config: ModelConfig) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None:
|
||||||
|
# No-op: all parameters are already on meta from init.
|
||||||
|
pass
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
self,
|
||||||
|
vllm_config: VllmConfig,
|
||||||
|
model_config: ModelConfig,
|
||||||
|
prefix: str = "",
|
||||||
|
) -> nn.Module:
|
||||||
|
device_config = vllm_config.device_config
|
||||||
|
load_config = vllm_config.load_config
|
||||||
|
load_device = (
|
||||||
|
device_config.device if load_config.device is None else load_config.device
|
||||||
|
)
|
||||||
|
target_device = torch.device(load_device)
|
||||||
|
|
||||||
|
with set_default_torch_dtype(model_config.dtype):
|
||||||
|
# Initialize model on meta device — no GPU memory, and
|
||||||
|
# Parameter subclasses are preserved (unlike FakeTensorMode
|
||||||
|
# which converts them to FakeTensor).
|
||||||
|
with torch.device("meta"):
|
||||||
|
model = initialize_model(
|
||||||
|
vllm_config=vllm_config,
|
||||||
|
model_config=model_config,
|
||||||
|
prefix=prefix,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Run weight post-processing on meta tensors.
|
||||||
|
from vllm.model_executor.model_loader.base_loader import (
|
||||||
|
_has_online_quant,
|
||||||
|
)
|
||||||
|
from vllm.model_executor.model_loader.reload import (
|
||||||
|
finalize_layerwise_processing,
|
||||||
|
)
|
||||||
|
|
||||||
|
if _has_online_quant(model):
|
||||||
|
finalize_layerwise_processing(model, model_config)
|
||||||
|
process_weights_after_loading(model, model_config, target_device)
|
||||||
|
|
||||||
|
# Set fake_device on all parameters and buffers so
|
||||||
|
# swap_meta_params_to_fake knows the intended target device.
|
||||||
|
# Must happen AFTER process_weights_after_loading which may
|
||||||
|
# create new parameters.
|
||||||
|
for param in model.parameters():
|
||||||
|
param.fake_device = target_device
|
||||||
|
for buf in model.buffers():
|
||||||
|
buf.fake_device = target_device
|
||||||
|
|
||||||
|
return model.eval()
|
||||||
|
|
||||||
|
|
||||||
|
def swap_meta_params_to_fake(model: nn.Module) -> None:
|
||||||
|
"""Replace all meta-device parameters and buffers with FakeTensors.
|
||||||
|
|
||||||
|
Called before torch.compile so that Dynamo sees cuda-device tensors
|
||||||
|
during tracing. Parameter subclasses are replaced with plain
|
||||||
|
FakeTensors at this point — isinstance checks are no longer needed
|
||||||
|
after weight processing is complete.
|
||||||
|
"""
|
||||||
|
from torch._subclasses.fake_tensor import FakeTensorMode
|
||||||
|
|
||||||
|
fake_mode = FakeTensorMode()
|
||||||
|
with fake_mode:
|
||||||
|
for name, param in list(model.named_parameters()):
|
||||||
|
device = getattr(param, "fake_device", None)
|
||||||
|
assert device is not None and device != torch.device("meta"), (
|
||||||
|
f"Parameter {name} missing fake_device or has meta device"
|
||||||
|
)
|
||||||
|
fake_data = torch.empty(
|
||||||
|
param.shape,
|
||||||
|
dtype=param.dtype,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
# Navigate to the owning module and replace the parameter.
|
||||||
|
*path, attr = name.split(".")
|
||||||
|
parent = model.get_submodule(".".join(path)) if path else model
|
||||||
|
parent.register_parameter(
|
||||||
|
attr,
|
||||||
|
nn.Parameter(fake_data, requires_grad=param.requires_grad),
|
||||||
|
)
|
||||||
|
for name, buf in list(model.named_buffers()):
|
||||||
|
device = getattr(buf, "fake_device", buf.device)
|
||||||
|
if device == torch.device("meta"):
|
||||||
|
device = torch.device("cuda")
|
||||||
|
fake_buf = torch.empty(
|
||||||
|
buf.shape,
|
||||||
|
dtype=buf.dtype,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
*path, attr = name.split(".")
|
||||||
|
parent = model.get_submodule(".".join(path)) if path else model
|
||||||
|
parent.register_buffer(attr, fake_buf)
|
||||||
@@ -4763,6 +4763,12 @@ class GPUModelRunner(
|
|||||||
if load_dummy_weights:
|
if load_dummy_weights:
|
||||||
self.load_config.load_format = "dummy"
|
self.load_config.load_format = "dummy"
|
||||||
model_loader = get_model_loader(self.load_config)
|
model_loader = get_model_loader(self.load_config)
|
||||||
|
if self.vllm_config.compilation_config.compile_only:
|
||||||
|
from vllm.model_executor.model_loader.fake_loader import (
|
||||||
|
FakeModelLoader,
|
||||||
|
)
|
||||||
|
|
||||||
|
model_loader = FakeModelLoader()
|
||||||
self.model = model_loader.load_model(
|
self.model = model_loader.load_model(
|
||||||
vllm_config=self.vllm_config, model_config=self.model_config
|
vllm_config=self.vllm_config, model_config=self.model_config
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -5,7 +5,7 @@
|
|||||||
import gc
|
import gc
|
||||||
import os
|
import os
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from contextlib import AbstractContextManager, nullcontext
|
from contextlib import AbstractContextManager, nullcontext, suppress
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
from types import NoneType
|
from types import NoneType
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
@@ -280,7 +280,12 @@ class Worker(WorkerBase):
|
|||||||
|
|
||||||
# take current memory snapshot
|
# take current memory snapshot
|
||||||
self.init_snapshot = init_snapshot = MemorySnapshot(device=self.device)
|
self.init_snapshot = init_snapshot = MemorySnapshot(device=self.device)
|
||||||
self.requested_memory = request_memory(init_snapshot, self.cache_config)
|
if self.compilation_config.compile_only:
|
||||||
|
# In compile-only mode with fake weights, skip memory
|
||||||
|
# validation since we don't allocate real GPU memory.
|
||||||
|
self.requested_memory = 0
|
||||||
|
else:
|
||||||
|
self.requested_memory = request_memory(init_snapshot, self.cache_config)
|
||||||
logger.debug("worker init memory snapshot: %r", self.init_snapshot)
|
logger.debug("worker init memory snapshot: %r", self.init_snapshot)
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"worker requested memory: %sGiB", format_gib(self.requested_memory)
|
"worker requested memory: %sGiB", format_gib(self.requested_memory)
|
||||||
@@ -573,10 +578,59 @@ class Worker(WorkerBase):
|
|||||||
if not any(x in compile_range for x in all_sizes):
|
if not any(x in compile_range for x in all_sizes):
|
||||||
warmup_sizes.append(compile_range.end)
|
warmup_sizes.append(compile_range.end)
|
||||||
|
|
||||||
|
if self.compilation_config.compile_only:
|
||||||
|
from vllm.compilation.backends import CompilationDone
|
||||||
|
from vllm.model_executor.model_loader.fake_loader import (
|
||||||
|
swap_meta_params_to_fake,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Swap meta-device parameters to FakeTensors so that
|
||||||
|
# torch.compile sees cuda-device tensors during tracing.
|
||||||
|
swap_meta_params_to_fake(self.model_runner.model)
|
||||||
|
|
||||||
|
# Verify that no significant GPU memory was allocated for
|
||||||
|
# model weights. A small amount (< 64 MiB) may come from
|
||||||
|
# CUDA runtime or library initialization.
|
||||||
|
_COMPILE_ONLY_MEM_THRESHOLD = 64 * 1024 * 1024 # 64 MiB
|
||||||
|
mem_used = torch.accelerator.memory_allocated(self.device)
|
||||||
|
assert mem_used < _COMPILE_ONLY_MEM_THRESHOLD, (
|
||||||
|
f"compile-only mode should use minimal GPU memory after "
|
||||||
|
f"model loading, but {format_gib(mem_used)} GiB is "
|
||||||
|
f"allocated (threshold: "
|
||||||
|
f"{format_gib(_COMPILE_ONLY_MEM_THRESHOLD)} GiB)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# In the normal path, the first torch.compile is triggered
|
||||||
|
# by profile_run() which calls _dummy_run(max_num_tokens).
|
||||||
|
# In compile-only mode we skip _initialize_kv_caches (which
|
||||||
|
# calls profile_run), so call it here to trigger compilation.
|
||||||
|
# CompilationDone is raised after vLLM's torch.compile cache
|
||||||
|
# and AOT artifact are saved, to prevent execution with fake
|
||||||
|
# tensors.
|
||||||
|
with suppress(CompilationDone):
|
||||||
|
self.model_runner.profile_run()
|
||||||
|
|
||||||
|
# Verify compilation didn't allocate significant GPU memory.
|
||||||
|
mem_used = torch.accelerator.memory_allocated(self.device)
|
||||||
|
assert mem_used < _COMPILE_ONLY_MEM_THRESHOLD, (
|
||||||
|
f"compile-only mode should use minimal GPU memory after "
|
||||||
|
f"compilation, but {format_gib(mem_used)} GiB is "
|
||||||
|
f"allocated (threshold: "
|
||||||
|
f"{format_gib(_COMPILE_ONLY_MEM_THRESHOLD)} GiB)"
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Compile-only mode: compilation complete. "
|
||||||
|
"Skipping kernel warmup, CUDA graphs, and "
|
||||||
|
"sampler warmup."
|
||||||
|
)
|
||||||
|
return self.compilation_config.compilation_time
|
||||||
|
|
||||||
# We skip EPLB here since we don't want to record dummy metrics
|
# We skip EPLB here since we don't want to record dummy metrics
|
||||||
for size in sorted(warmup_sizes, reverse=True):
|
for size in sorted(warmup_sizes, reverse=True):
|
||||||
logger.info("Compile and warming up model for size %d", size)
|
logger.info("Compile and warming up model for size %d", size)
|
||||||
self.model_runner._dummy_run(size, skip_eplb=True, remove_lora=False)
|
self.model_runner._dummy_run(size, skip_eplb=True, remove_lora=False)
|
||||||
|
|
||||||
self.model_runner.maybe_remove_all_loras(self.model_runner.lora_config)
|
self.model_runner.maybe_remove_all_loras(self.model_runner.lora_config)
|
||||||
|
|
||||||
# Warmup and tune the kernels used during model execution before
|
# Warmup and tune the kernels used during model execution before
|
||||||
|
|||||||
Reference in New Issue
Block a user