diff --git a/docs/serving/offline_inference.md b/docs/serving/offline_inference.md index 4512f4a0720..9a71612f262 100644 --- a/docs/serving/offline_inference.md +++ b/docs/serving/offline_inference.md @@ -65,6 +65,8 @@ For further details on Weight Transfer, please refer to [this page](../training/ - `LLM.start_weight_update` - Starts a new weight update cycle. - `LLM.update_weights` - Updates the model weights. - `LLM.finish_weight_update` - Finishes the current weight update cycle. +- `LLM.update_weight_version` - Sets the weight version without updating model weights. +- `LLM.get_weight_version` - Returns the latest committed weight version. ## Additional APIs diff --git a/docs/serving/online_serving/README.md b/docs/serving/online_serving/README.md index f7914bc0582..90ff7a3e3d8 100644 --- a/docs/serving/online_serving/README.md +++ b/docs/serving/online_serving/README.md @@ -179,6 +179,8 @@ For further details on Weight Transfer, please refer to [this page](../../traini - `/start_weight_update` - Prepares the inference engine for a weight update. - `/update_weights` - Update model weights (can alter model behavior) - `/finish_weight_update` - Finalizes the weight update +- `/update_weight_version` - Set the weight version without updating model weights +- `/weight_info` - Get the latest committed weight version - `/get_world_size` - Get distributed world size ### Collective RPC diff --git a/docs/training/async_rl.md b/docs/training/async_rl.md index e655f9c39ff..9e75a24eaa1 100644 --- a/docs/training/async_rl.md +++ b/docs/training/async_rl.md @@ -38,11 +38,12 @@ Resumes the scheduler after a pause. Any requests frozen with `mode="keep"` will ### HTTP Endpoints -When using the vLLM HTTP server, the same functionality is available via: +With `VLLM_SERVER_DEV_MODE=1`, the vLLM HTTP server exposes the same functionality via: - `POST /pause?mode=keep` - Pause generation - `POST /resume` - Resume generation - `POST /abort_requests` - Abort in-flight requests without pausing the scheduler (send `{}` to abort all, or `{"request_ids": [...]}`) +- `GET /weight_info` - Return the latest committed `weight_version` !!! note "Data Parallelism" When using data parallelism with vLLM's **internal load balancer** (i.e. `data_parallel_backend="ray"`), pause and resume are handled automatically across all DP ranks -- a single call is sufficient. When using an **external load balancer** (i.e. multiple independent vLLM instances behind a proxy), you must send pause and resume requests to **every** engine instance individually before and after the weight update. diff --git a/docs/training/weight_transfer/README.md b/docs/training/weight_transfer/README.md index 7579e5fd4d0..b8d39763181 100644 --- a/docs/training/weight_transfer/README.md +++ b/docs/training/weight_transfer/README.md @@ -53,7 +53,9 @@ When running vLLM as an HTTP server, the following endpoints are available for w | `/init_weight_transfer_engine` | POST | Initialize the weight transfer engine with backend-specific info | | `/start_weight_update` | POST | Start a weight update | | `/update_weights` | POST | Transfer a batch of weights with backend-specific metadata | -| `/finish_weight_update` | POST | Finish the weight update and run post-processing | +| `/finish_weight_update` | POST | Finish the update and optionally commit its `weight_version` | +| `/update_weight_version` | POST | Update `weight_version` without changing model weights | +| `/weight_info` | GET | Get the latest committed weight version | | `/pause` | POST | Pause generation before weight sync to handle inflight requests | | `/resume` | POST | Resume generation after weight sync | | `/get_world_size` | GET | Get the number of inference workers (useful for NCCL world size calculation) | @@ -79,7 +81,7 @@ EngineClass.trainer_send_weights( ) # 4. Finish weight update on inference side -llm.finish_weight_update() +llm.finish_weight_update(weight_version="step-42") ``` See the [NCCL](nccl.md) and [IPC](ipc.md) pages for backend-specific trainer APIs and full examples. diff --git a/tests/distributed/test_weight_transfer.py b/tests/distributed/test_weight_transfer.py index b79aa1974d1..eeeceb95998 100644 --- a/tests/distributed/test_weight_transfer.py +++ b/tests/distributed/test_weight_transfer.py @@ -1247,7 +1247,7 @@ class RecordingClient: self.order.append("update") self.last_update_info = update_info - def finish_weight_update(self) -> None: + def finish_weight_update(self, weight_version: str | None = None) -> None: self.order.append("finish") @@ -1303,6 +1303,10 @@ class TestTrainerClients: assert isinstance(update_req, WeightTransferUpdateRequest) assert update_req.update_info == {"names": ["w"]} + client.finish_weight_update("step-42") + handle.finish_weight_update.remote.assert_called_once_with() + handle.update_weight_version.remote.assert_called_once_with("step-42") + def test_http_client_pickles_ipc_handles_for_json(self, monkeypatch): """HTTP update_weights must encode raw ipc_handles as a base64 pickle.""" captured = {} @@ -1334,6 +1338,9 @@ class TestTrainerClients: client.update_weights(update_info) assert captured["json"]["update_info"] == update_info + client.finish_weight_update("step-42") + assert captured["json"] == {"weight_version": "step-42"} + class TestModuleSource: """`ModuleSource` metadata vs. materialized iteration (dense, no GPU).""" diff --git a/tests/entrypoints/openai/test_openai_schema.py b/tests/entrypoints/openai/test_openai_schema.py index 2985c539518..6d3fc2f4474 100644 --- a/tests/entrypoints/openai/test_openai_schema.py +++ b/tests/entrypoints/openai/test_openai_schema.py @@ -148,6 +148,7 @@ def test_openapi_stateless(case: schemathesis.Case): "/start_draft_weight_update", "/update_weights", "/finish_weight_update", + "/update_weight_version", ): return diff --git a/tests/entrypoints/weight_transfer/test_weight_transfer_llm.py b/tests/entrypoints/weight_transfer/test_weight_transfer_llm.py index 9088b3c5e8d..31d562e5ccf 100644 --- a/tests/entrypoints/weight_transfer/test_weight_transfer_llm.py +++ b/tests/entrypoints/weight_transfer/test_weight_transfer_llm.py @@ -234,6 +234,7 @@ def test_update_weights_calls_engine(): assert shapes == test_shapes llm.finish_weight_update() + assert llm.get_weight_version() == "default" @create_new_process_for_each_test() @@ -259,6 +260,8 @@ def test_full_weight_transfer_flow(): weight_transfer_config=WeightTransferConfig(backend="nccl"), ) + assert llm.get_weight_version() == "default" + # Step 1: Initialize weight transfer engine llm.init_weight_transfer_engine( WeightTransferInitRequest(init_info={"test_param": "flow_test"}) @@ -278,8 +281,15 @@ def test_full_weight_transfer_flow(): ) ) + assert llm.get_weight_version() == "default" + # Step 4: Finish weight update - llm.finish_weight_update() + llm.finish_weight_update("step-42") + + assert llm.get_weight_version() == "step-42" + + llm.update_weight_version("manual-version") + assert llm.get_weight_version() == "manual-version" # Verify the full flow completed def check_flow(self): diff --git a/vllm/distributed/weight_transfer/base.py b/vllm/distributed/weight_transfer/base.py index 2e377e29253..adddf41ff4e 100644 --- a/vllm/distributed/weight_transfer/base.py +++ b/vllm/distributed/weight_transfer/base.py @@ -370,7 +370,7 @@ class VLLMWeightSyncClient(Protocol): def update_weights(self, update_info: dict[str, Any]) -> None: ... - def finish_weight_update(self) -> None: ... + def finish_weight_update(self, weight_version: str | None = None) -> None: ... class TrainerWeightTransferEngine(ABC, Generic[TConfig, TInitInfo]): diff --git a/vllm/distributed/weight_transfer/clients.py b/vllm/distributed/weight_transfer/clients.py index 4f54a6e291e..12dd0c9eacc 100644 --- a/vllm/distributed/weight_transfer/clients.py +++ b/vllm/distributed/weight_transfer/clients.py @@ -77,8 +77,11 @@ class HTTPVLLMWeightSyncClient: "update_weights", {"update_info": _json_safe_update_info(update_info)} ) - def finish_weight_update(self) -> None: - self._post("finish_weight_update") + def finish_weight_update(self, weight_version: str | None = None) -> None: + json = ( + {"weight_version": weight_version} if weight_version is not None else None + ) + self._post("finish_weight_update", json) class RayVLLMWeightSyncClient: @@ -108,7 +111,11 @@ class RayVLLMWeightSyncClient: request = WeightTransferUpdateRequest(update_info=update_info) ray.get([h.update_weights.remote(request) for h in self.handles]) - def finish_weight_update(self) -> None: + def finish_weight_update(self, weight_version: str | None = None) -> None: import ray ray.get([h.finish_weight_update.remote() for h in self.handles]) + if weight_version is not None: + ray.get( + [h.update_weight_version.remote(weight_version) for h in self.handles] + ) diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py index ef3be178ac8..5a9b9f96d2c 100644 --- a/vllm/engine/protocol.py +++ b/vllm/engine/protocol.py @@ -267,6 +267,14 @@ class EngineClient(ABC): """Batched weight update for RL training.""" raise NotImplementedError - async def finish_weight_update(self) -> None: - """Finish the current weight update.""" + async def finish_weight_update(self, weight_version: str | None = None) -> None: + """Finish the weight update and set its version if provided.""" + raise NotImplementedError + + async def update_weight_version(self, new_version: str) -> None: + """Set the weight version without updating weights.""" + raise NotImplementedError + + async def get_weight_version(self) -> str: + """Return the latest committed weight version.""" raise NotImplementedError diff --git a/vllm/entrypoints/llm.py b/vllm/entrypoints/llm.py index b3205728e49..4274819d988 100644 --- a/vllm/entrypoints/llm.py +++ b/vllm/entrypoints/llm.py @@ -885,9 +885,19 @@ class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin, OfflineInferenceMixin): "update_weights", kwargs={"update_info": update_info_dict} ) - def finish_weight_update(self) -> None: - """Finish the current weight update.""" + def finish_weight_update(self, weight_version: str | None = None) -> None: + """Finish the weight update and set its version if provided.""" self.llm_engine.collective_rpc("finish_weight_update") + if weight_version is not None: + self.llm_engine.set_weight_version(weight_version) + + def update_weight_version(self, new_version: str) -> None: + """Set the weight version without updating weights.""" + self.llm_engine.set_weight_version(new_version) + + def get_weight_version(self) -> str: + """Return the latest committed weight version.""" + return self.llm_engine.get_weight_version() def __repr__(self) -> str: """Return a transformers-style hierarchical view of the model.""" diff --git a/vllm/entrypoints/serve/dev/rlhf/api_router.py b/vllm/entrypoints/serve/dev/rlhf/api_router.py index 8a2494a59df..392fcf56747 100644 --- a/vllm/entrypoints/serve/dev/rlhf/api_router.py +++ b/vllm/entrypoints/serve/dev/rlhf/api_router.py @@ -5,7 +5,7 @@ import json from http import HTTPStatus from typing import Annotated -from fastapi import APIRouter, FastAPI, HTTPException, Query, Request +from fastapi import APIRouter, Body, FastAPI, HTTPException, Query, Request from fastapi.responses import JSONResponse from vllm.distributed.weight_transfer.base import ( @@ -203,11 +203,29 @@ async def update_weights(raw_request: Request): @router.post("/finish_weight_update") -async def finish_weight_update(raw_request: Request): - await engine_client(raw_request).finish_weight_update() +async def finish_weight_update( + raw_request: Request, + weight_version: Annotated[str | None, Body(embed=True)] = None, +): + await engine_client(raw_request).finish_weight_update(weight_version) return JSONResponse(content={"message": "Weight update finished"}) +@router.post("/update_weight_version") +async def update_weight_version( + raw_request: Request, + new_version: Annotated[str, Body(embed=True)], +): + await engine_client(raw_request).update_weight_version(new_version) + return JSONResponse(content={"success": True, "new_version": new_version}) + + +@router.get("/weight_info") +async def weight_info(raw_request: Request): + weight_version = await engine_client(raw_request).get_weight_version() + return JSONResponse(content={"weight_version": weight_version}) + + @router.get("/get_world_size") async def get_world_size( raw_request: Request, diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py index 922a8aa5982..5c2e01cf44b 100644 --- a/vllm/v1/engine/async_llm.py +++ b/vllm/v1/engine/async_llm.py @@ -1106,6 +1106,16 @@ class AsyncLLM(EngineClient): "update_weights", kwargs={"update_info": request.update_info} ) - async def finish_weight_update(self) -> None: - """Finish the current weight update.""" + async def finish_weight_update(self, weight_version: str | None = None) -> None: + """Finish the weight update and set its version if provided.""" await self.collective_rpc("finish_weight_update") + if weight_version is not None: + await self.update_weight_version(weight_version) + + async def update_weight_version(self, new_version: str) -> None: + """Set the weight version without updating weights.""" + await self.engine_core.set_weight_version_async(new_version) + + async def get_weight_version(self) -> str: + """Return the latest committed weight version.""" + return await self.engine_core.get_weight_version_async() diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index 9817c474343..9917f810b5b 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -125,6 +125,8 @@ class EngineCore: ) self.log_stats = log_stats + # Opaque weight version supplied by the caller. + self._weight_version = "default" # Setup Model. self.model_executor = executor_class(vllm_config) @@ -956,6 +958,13 @@ class EngineCore: ) -> list[_R]: return self.model_executor.collective_rpc(method, timeout, args, kwargs) + def set_weight_version(self, weight_version: str) -> None: + self._weight_version = weight_version + + def get_weight_version(self) -> str: + """Return the latest committed weight version.""" + return self._weight_version + def preprocess_add_request(self, request: EngineCoreRequest) -> tuple[Request, int]: """Preprocess the request. diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index 0aa4b6f3312..a6c232b2ab7 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -176,9 +176,21 @@ class EngineCoreClient(ABC): def execute_dummy_batch(self) -> None: raise NotImplementedError + def set_weight_version(self, weight_version: str) -> None: + raise NotImplementedError + + def get_weight_version(self) -> str: + raise NotImplementedError + async def execute_dummy_batch_async(self) -> None: raise NotImplementedError + async def set_weight_version_async(self, weight_version: str) -> None: + raise NotImplementedError + + async def get_weight_version_async(self) -> str: + raise NotImplementedError + def abort_requests(self, request_ids: list[str]) -> None: raise NotImplementedError @@ -351,6 +363,12 @@ class InprocClient(EngineCoreClient): def execute_dummy_batch(self) -> None: self.engine_core.execute_dummy_batch() + def set_weight_version(self, weight_version: str) -> None: + self.engine_core.set_weight_version(weight_version) + + def get_weight_version(self) -> str: + return self.engine_core.get_weight_version() + def add_lora(self, lora_request: LoRARequest) -> bool: return self.engine_core.add_lora(lora_request) @@ -947,6 +965,12 @@ class SyncMPClient(MPClient): def execute_dummy_batch(self) -> None: self.call_utility("execute_dummy_batch") + def set_weight_version(self, weight_version: str) -> None: + self.call_utility("set_weight_version", weight_version) + + def get_weight_version(self) -> str: + return self.call_utility("get_weight_version") + def collective_rpc( self, method: str | Callable[..., _R], @@ -1199,6 +1223,12 @@ class AsyncMPClient(MPClient): async def execute_dummy_batch_async(self) -> None: await self.call_utility_async("execute_dummy_batch") + async def set_weight_version_async(self, weight_version: str) -> None: + await self.call_utility_async("set_weight_version", weight_version) + + async def get_weight_version_async(self) -> str: + return await self.call_utility_async("get_weight_version") + async def add_lora_async(self, lora_request: LoRARequest) -> bool: return await self.call_utility_async("add_lora", lora_request) diff --git a/vllm/v1/engine/llm_engine.py b/vllm/v1/engine/llm_engine.py index ff86a1dffd9..17e40630859 100644 --- a/vllm/v1/engine/llm_engine.py +++ b/vllm/v1/engine/llm_engine.py @@ -425,6 +425,13 @@ class LLMEngine: ) -> list[_R]: return self.engine_core.collective_rpc(method, timeout, args, kwargs) + def set_weight_version(self, weight_version: str) -> None: + self.engine_core.set_weight_version(weight_version) + + def get_weight_version(self) -> str: + """Return the latest committed weight version.""" + return self.engine_core.get_weight_version() + def apply_model(self, func: Callable[[nn.Module], _R]) -> list[_R]: return self.collective_rpc("apply_model", args=(func,))