[PD] Introduce runtime role switching between prefill and decode (#28403)
Signed-off-by: huanglong <huanglong@linux.alibaba.com> Signed-off-by: inkcherry <mingzhi.liu@amd.com> Co-authored-by: huanglong <huanglong@linux.alibaba.com> Co-authored-by: Shangming Cai <csmthu@gmail.com> Co-authored-by: Huang Long <121648372+LLLL114@users.noreply.github.com>
This commit is contained in:
co-authored by
huanglong
Shangming Cai
Huang Long
parent
a98d921658
commit
1f60ddef5d
@@ -111,6 +111,42 @@ class MiniLoadBalancer:
|
||||
self.decode_urls[didx],
|
||||
)
|
||||
|
||||
def current_role_and_port(self, worker_url):
|
||||
"""Return (role, bootstrap_port) of a registered server, or (None, None)
|
||||
if it is not in either routing list."""
|
||||
if worker_url in self.prefill_urls:
|
||||
return (
|
||||
"prefill",
|
||||
self.prefill_bootstrap_ports[self.prefill_urls.index(worker_url)],
|
||||
)
|
||||
if worker_url in self.decode_urls:
|
||||
return "decode", None
|
||||
return None, None
|
||||
|
||||
def remove_worker(self, worker_url):
|
||||
"""Drop a server from both routing lists so no new requests are sent to
|
||||
it (used to quiesce it before a role switch)."""
|
||||
if worker_url in self.decode_urls:
|
||||
self.decode_urls.remove(worker_url)
|
||||
if worker_url in self.prefill_urls:
|
||||
idx = self.prefill_urls.index(worker_url)
|
||||
self.prefill_urls.pop(idx)
|
||||
self.prefill_bootstrap_ports.pop(idx)
|
||||
|
||||
def add_worker(self, worker_url, role, bootstrap_port=None):
|
||||
"""Register a server under a role in the routing lists."""
|
||||
if role == "prefill":
|
||||
self.prefill_urls.append(worker_url)
|
||||
self.prefill_bootstrap_ports.append(bootstrap_port or 8998)
|
||||
elif role == "decode":
|
||||
self.decode_urls.append(worker_url)
|
||||
|
||||
def apply_role_switch(self, worker_url, new_role, bootstrap_port=None):
|
||||
"""Move a server between the prefill and decode routing lists after its
|
||||
role has been switched on the backend. Idempotent."""
|
||||
self.remove_worker(worker_url)
|
||||
self.add_worker(worker_url, new_role, bootstrap_port)
|
||||
|
||||
async def generate(
|
||||
self, modified_request, prefill_server, decode_server, endpoint
|
||||
) -> ORJSONResponse:
|
||||
@@ -255,6 +291,69 @@ async def health_generate():
|
||||
return Response(status_code=200)
|
||||
|
||||
|
||||
async def _post_role_switch(worker_url, body):
|
||||
"""POST the role switch to a backend server; return (status, json)."""
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
timeout=aiohttp.ClientTimeout(total=lb.timeout)
|
||||
) as session:
|
||||
async with session.post(f"{worker_url}/pd_role_switch", json=body) as resp:
|
||||
return resp.status, await resp.json()
|
||||
except Exception as e: # transport error -> report as a failure
|
||||
return 502, {"success": False, "message": str(e)}
|
||||
|
||||
|
||||
@app.post("/pd_role_switch")
|
||||
async def pd_role_switch(request_data: dict):
|
||||
"""Switch a running server's PD role (prefill<->decode) at runtime and
|
||||
update the LB's routing lists. Body: {"worker_url", "new_role":
|
||||
"prefill"|"decode", "bootstrap_port"?, "decode_cuda_graph_bs"?,
|
||||
"decode_cuda_graph_memory_gb"?, "drain"?, "drain_timeout_secs"?}.
|
||||
|
||||
The backend rejects a switch unless the instance is idle. To make this
|
||||
safe while serving, by default the LB first removes the server from its
|
||||
routing lists (so no new requests arrive), then retries the switch while
|
||||
the server drains its in-flight requests, and only then registers it
|
||||
under the new role. A failed server is restored only when the backend
|
||||
confirms that no role state changed."""
|
||||
worker_url = request_data.get("worker_url")
|
||||
new_role = request_data.get("new_role")
|
||||
if worker_url is None:
|
||||
raise HTTPException(status_code=400, detail="worker_url is required")
|
||||
if new_role not in ("prefill", "decode"):
|
||||
raise HTTPException(status_code=400, detail=f"invalid new_role={new_role!r}")
|
||||
|
||||
drain = request_data.get("drain", True)
|
||||
drain_timeout = request_data.get("drain_timeout_secs", 300)
|
||||
old_role, old_port = lb.current_role_and_port(worker_url)
|
||||
|
||||
body = {"new_role": new_role}
|
||||
for field in ("decode_cuda_graph_bs", "decode_cuda_graph_memory_gb"):
|
||||
if request_data.get(field) is not None:
|
||||
body[field] = request_data[field]
|
||||
|
||||
# Stop routing new requests to this server so it can drain to idle.
|
||||
if drain and old_role is not None:
|
||||
lb.remove_worker(worker_url)
|
||||
|
||||
deadline = asyncio.get_event_loop().time() + drain_timeout
|
||||
while True:
|
||||
status, result = await _post_role_switch(worker_url, body)
|
||||
if status == 200 and result.get("success", False):
|
||||
break
|
||||
# The backend rejects while not idle; keep retrying as it drains.
|
||||
not_idle = "not idle" in (result.get("message", "") or "").lower()
|
||||
if drain and not_idle and asyncio.get_event_loop().time() < deadline:
|
||||
await asyncio.sleep(1.0)
|
||||
continue
|
||||
if drain and old_role is not None and result.get("safe_to_restore", False):
|
||||
lb.add_worker(worker_url, old_role, old_port)
|
||||
return ORJSONResponse(content=result, status_code=status)
|
||||
|
||||
lb.apply_role_switch(worker_url, new_role, request_data.get("bootstrap_port"))
|
||||
return ORJSONResponse(content=result, status_code=200)
|
||||
|
||||
|
||||
@app.post("/flush_cache")
|
||||
async def flush_cache(timeout: Optional[float] = None):
|
||||
# `timeout` must reach the workers. The scheduler treats a missing or
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from sglang_router import mini_lb
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("result", "restored"),
|
||||
[
|
||||
(
|
||||
{
|
||||
"success": False,
|
||||
"message": "instance is not idle",
|
||||
"safe_to_restore": True,
|
||||
},
|
||||
True,
|
||||
),
|
||||
(
|
||||
{
|
||||
"success": False,
|
||||
"message": "instance unhealthy, restart required",
|
||||
"safe_to_restore": False,
|
||||
},
|
||||
False,
|
||||
),
|
||||
({"success": False, "message": "connection lost"}, False),
|
||||
],
|
||||
)
|
||||
def test_failed_role_switch_restores_only_healthy_worker(monkeypatch, result, restored):
|
||||
worker_url = "http://prefill:8000"
|
||||
load_balancer = mini_lb.MiniLoadBalancer.__new__(mini_lb.MiniLoadBalancer)
|
||||
load_balancer.timeout = 1
|
||||
load_balancer.prefill_urls = [worker_url]
|
||||
load_balancer.prefill_bootstrap_ports = [8998]
|
||||
load_balancer.decode_urls = ["http://decode:8000"]
|
||||
monkeypatch.setattr(mini_lb, "lb", load_balancer)
|
||||
|
||||
async def post_role_switch(*_args, **_kwargs):
|
||||
return 400, result
|
||||
|
||||
monkeypatch.setattr(mini_lb, "_post_role_switch", post_role_switch)
|
||||
response = asyncio.run(
|
||||
mini_lb.pd_role_switch(
|
||||
{
|
||||
"worker_url": worker_url,
|
||||
"new_role": "decode",
|
||||
"drain_timeout_secs": 0,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert (worker_url in load_balancer.prefill_urls) is restored
|
||||
|
||||
|
||||
def test_role_switch_forwards_decode_graph_requirements(monkeypatch):
|
||||
worker_url = "http://prefill:8000"
|
||||
load_balancer = mini_lb.MiniLoadBalancer.__new__(mini_lb.MiniLoadBalancer)
|
||||
load_balancer.timeout = 1
|
||||
load_balancer.prefill_urls = [worker_url]
|
||||
load_balancer.prefill_bootstrap_ports = [8998]
|
||||
load_balancer.decode_urls = ["http://decode:8000"]
|
||||
monkeypatch.setattr(mini_lb, "lb", load_balancer)
|
||||
sent_body = {}
|
||||
|
||||
async def post_role_switch(_worker_url, body):
|
||||
sent_body.update(body)
|
||||
return 200, {"success": True, "message": "ok"}
|
||||
|
||||
monkeypatch.setattr(mini_lb, "_post_role_switch", post_role_switch)
|
||||
response = asyncio.run(
|
||||
mini_lb.pd_role_switch(
|
||||
{
|
||||
"worker_url": worker_url,
|
||||
"new_role": "decode",
|
||||
"decode_cuda_graph_bs": [1, 2, 4],
|
||||
"decode_cuda_graph_memory_gb": 1.25,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert sent_body["decode_cuda_graph_bs"] == [1, 2, 4]
|
||||
assert sent_body["decode_cuda_graph_memory_gb"] == 1.25
|
||||
Reference in New Issue
Block a user