From 802a4fce3ebc823e5a185c20863ccd919e2ca105 Mon Sep 17 00:00:00 2001 From: yewentao256 Date: Fri, 7 Aug 2026 15:30:39 +0000 Subject: [PATCH] MR v2 weight offloading support Signed-off-by: yewentao256 --- tests/basic_correctness/test_cpu_offload.py | 6 ++++-- tests/basic_correctness/test_prefetch_offload.py | 7 ++++++- vllm/v1/worker/gpu/model_runner.py | 9 +++++++++ 3 files changed, 19 insertions(+), 3 deletions(-) diff --git a/tests/basic_correctness/test_cpu_offload.py b/tests/basic_correctness/test_cpu_offload.py index c1df36b369a9..c9ca34281854 100644 --- a/tests/basic_correctness/test_cpu_offload.py +++ b/tests/basic_correctness/test_cpu_offload.py @@ -8,8 +8,10 @@ @pytest.mark.parametrize("disable_pin_memory", [False, True]) @pytest.mark.parametrize("disable_uva", [False, True]) -def test_cpu_offload(disable_pin_memory, disable_uva): +@pytest.mark.parametrize("use_v2_model_runner", [False, True]) +def test_cpu_offload(disable_pin_memory, disable_uva, use_v2_model_runner): env_vars = { + "VLLM_USE_V2_MODEL_RUNNER": str(int(use_v2_model_runner)), "VLLM_WEIGHT_OFFLOADING_DISABLE_PIN_MEMORY": str(int(disable_pin_memory)), "VLLM_WEIGHT_OFFLOADING_DISABLE_UVA": str(int(disable_uva)), } @@ -24,6 +26,6 @@ def test_cpu_offload(disable_pin_memory, disable_uva): model="hmellor/tiny-random-LlamaForCausalLM", arg1=[], arg2=args, - env1=None, + env1=env_vars, env2=env_vars, ) diff --git a/tests/basic_correctness/test_prefetch_offload.py b/tests/basic_correctness/test_prefetch_offload.py index 498887024ee6..125d36418873 100644 --- a/tests/basic_correctness/test_prefetch_offload.py +++ b/tests/basic_correctness/test_prefetch_offload.py @@ -2,10 +2,13 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """Test prefetch offloading correctness with Llama model.""" +import pytest + from ..utils import compare_two_settings -def test_prefetch_offload_llama(): +@pytest.mark.parametrize("use_v2_model_runner", [False, True]) +def test_prefetch_offload_llama(use_v2_model_runner): """Test prefetch CPU offloading with Llama-3.2-1B-Instruct. Compares outputs between: @@ -30,4 +33,6 @@ def test_prefetch_offload_llama(): "down_proj", ], [], # Baseline: no offloading + env1={"VLLM_USE_V2_MODEL_RUNNER": str(int(use_v2_model_runner))}, + env2={"VLLM_USE_V2_MODEL_RUNNER": str(int(use_v2_model_runner))}, ) diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index f245ecff23a4..a6ac6eb67e34 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -46,6 +46,11 @@ initialize_mamba_ssu_backend, ) from vllm.model_executor.model_loader import get_model_loader +from vllm.model_executor.offloader import ( + create_offloader, + get_offloader, + set_offloader, +) from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.encoder_budget import ( MultiModalBudget, @@ -292,6 +297,8 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.eplb = EPLBController(self.parallel_config, self.device) self.routed_experts_capturer: RoutedExpertsCapturer | None = None + set_offloader(create_offloader(self.vllm_config.offload_config)) + def update_max_model_len(self, max_model_len: int) -> None: self.max_model_len = max_model_len self.req_states.max_model_len = max_model_len @@ -413,6 +420,8 @@ def load_model(self, load_dummy_weights: bool = False, *args, **kwargs) -> None: device=self.device, ) + get_offloader().post_init() + def get_model(self) -> nn.Module: return self.model