[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:
Shuolei Wang
2026-07-28 14:35:29 +08:00
committed by GitHub
parent 90245f4190
commit 9069a57139
16 changed files with 142 additions and 18 deletions
+2
View File
@@ -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
+2
View File
@@ -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
+2 -1
View File
@@ -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.
+4 -2
View File
@@ -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.
+8 -1
View File
@@ -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):
+1 -1
View File
@@ -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]):
+10 -3
View File
@@ -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
View File
@@ -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
View File
@@ -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."""
+21 -3
View File
@@ -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,
+12 -2
View File
@@ -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()
+9
View File
@@ -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.
+30
View File
@@ -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)
+7
View File
@@ -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,))