forked from Karylab-cklius/vllm
[Core][Frontend] Add weight version tagging for RL rollouts (#49040)
Signed-off-by: Shuolei Wang <shuoleiwang123@gmail.com> Signed-off-by: Shuolei Wang <948904026@qq.com>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)."""
|
||||
|
||||
@@ -148,6 +148,7 @@ def test_openapi_stateless(case: schemathesis.Case):
|
||||
"/start_draft_weight_update",
|
||||
"/update_weights",
|
||||
"/finish_weight_update",
|
||||
"/update_weight_version",
|
||||
):
|
||||
return
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -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]
|
||||
)
|
||||
|
||||
+10
-2
@@ -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
|
||||
|
||||
+12
-2
@@ -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."""
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user