[Bugfix][RL] Set vLLM config during weight reload (#45989)

Signed-off-by: aoshen02 <aoshen@inferact.ai>
This commit is contained in:
aoshen02
2026-07-20 15:02:56 +08:00
committed by GitHub
parent 5245c80564
commit 9459fc6471
2 changed files with 45 additions and 12 deletions
@@ -9,6 +9,7 @@ session is active. These tests verify that delegation and the session guard.
import pytest
from vllm.config import VllmConfig, get_current_vllm_config
from vllm.v1.worker.gpu_worker import Worker
@@ -21,29 +22,55 @@ class _RecordingEngine:
self.finished = False
self.reset_count = 0
self.update_calls: list[dict] = []
self.seen_configs: list[VllmConfig] = []
def _record_config(self) -> None:
self.seen_configs.append(get_current_vllm_config())
def start_weight_update(self) -> None:
self._record_config()
self.started = True
def update_weights(self, update_info: dict) -> None:
self._record_config()
self.update_calls.append(update_info)
if self.raise_on_update:
raise ValueError("boom")
def finish_weight_update(self) -> None:
self._record_config()
self.finished = True
def reset_weight_update_target(self) -> None:
self.reset_count += 1
class _RecordingModelRunner:
def __init__(self) -> None:
self.seen_config: VllmConfig | None = None
def reload_weights(self) -> None:
self.seen_config = get_current_vllm_config()
def _make_worker(engine: _RecordingEngine | None) -> Worker:
worker = object.__new__(Worker)
worker.vllm_config = VllmConfig()
worker.weight_transfer_engine = engine
worker._weight_update_active = False
return worker
def test_reload_weights_sets_current_config():
worker = _make_worker(None)
model_runner = _RecordingModelRunner()
worker.model_runner = model_runner # type: ignore[assignment]
Worker.reload_weights(worker)
assert model_runner.seen_config is worker.vllm_config
def test_start_update_finish_delegates_to_engine():
engine = _RecordingEngine()
worker = _make_worker(engine)
@@ -60,6 +87,7 @@ def test_start_update_finish_delegates_to_engine():
assert engine.finished is True
assert engine.reset_count == 1
assert worker._weight_update_active is False
assert engine.seen_configs == [worker.vllm_config] * 3
def test_double_start_raises():
+17 -12
View File
@@ -442,7 +442,8 @@ class Worker(WorkerBase):
self.model_runner.update_config(overrides)
def reload_weights(self, *args, **kwargs) -> None:
self.model_runner.reload_weights(*args, **kwargs)
with set_current_vllm_config(self.vllm_config):
self.model_runner.reload_weights(*args, **kwargs)
@torch.inference_mode()
def determine_available_memory(self) -> int:
@@ -1301,14 +1302,16 @@ class Worker(WorkerBase):
the configured weight transfer engine. The worker only tracks that a
session is active.
"""
self._start_weight_update()
with set_current_vllm_config(self.vllm_config):
self._start_weight_update()
def start_draft_weight_update(self) -> None:
"""
Like start_weight_update, but retargets the engine at the speculative
draft model for this session.
"""
self._start_weight_update(is_draft=True)
with set_current_vllm_config(self.vllm_config):
self._start_weight_update(is_draft=True)
def _start_weight_update(self, is_draft: bool = False) -> None:
self._check_weight_transfer_engine()
@@ -1355,12 +1358,13 @@ class Worker(WorkerBase):
"start_weight_update must be called before update_weights."
)
try:
self.weight_transfer_engine.update_weights(update_info)
except BaseException:
self._weight_update_active = False
self.weight_transfer_engine.reset_weight_update_target()
raise
with set_current_vllm_config(self.vllm_config):
try:
self.weight_transfer_engine.update_weights(update_info)
except BaseException:
self._weight_update_active = False
self.weight_transfer_engine.reset_weight_update_target()
raise
def finish_weight_update(self) -> None:
"""Finish the current weight update session."""
@@ -1372,9 +1376,10 @@ class Worker(WorkerBase):
"finish_weight_update called without a matching start_weight_update."
)
self.weight_transfer_engine.finish_weight_update()
self.weight_transfer_engine.reset_weight_update_target()
self._weight_update_active = False
with set_current_vllm_config(self.vllm_config):
self.weight_transfer_engine.finish_weight_update()
self.weight_transfer_engine.reset_weight_update_target()
self._weight_update_active = False
def shutdown(self) -> None:
gc.unfreeze()