[PD-Disagg] Fully support external DP dispatch w/ PD-disaggregation mode. (#19268)

Co-authored-by: Ratish P <114130421+ratish1@users.noreply.github.com>
This commit is contained in:
Liangsheng Yin
2026-02-24 19:58:01 -08:00
committed by GitHub
co-authored by Ratish P
parent 241ee90164
commit 539f772f54
18 changed files with 253 additions and 62 deletions
@@ -7,6 +7,7 @@ import ipaddress
import logging
import random
import urllib
import warnings
from http import HTTPStatus
from itertools import chain
from typing import Optional
@@ -69,6 +70,10 @@ class MiniLoadBalancer:
)
self.enable_trace = False
self.test_external_dp_routing = router_args.test_external_dp_routing
self.prefill_dp_size = None
self.decode_dp_size = None
def _validate_router_args(self, router_args: RouterArgs):
logger.warning(
"\x1b[33mMiniLB is only for debugging purposes, it only supports random policy!\033[0m"
@@ -95,6 +100,32 @@ class MiniLoadBalancer:
trace_set_thread_info("Mini lb")
uvicorn.run(app, host=self.host, port=self.port)
async def _ensure_dp_sizes(self):
if self.prefill_dp_size is not None:
return
async with aiohttp.ClientSession() as session:
async with session.get(f"{self.prefill_urls[0]}/server_info") as resp:
info = await resp.json()
self.prefill_dp_size = len(info.get("internal_states", [1]))
async with session.get(f"{self.decode_urls[0]}/server_info") as resp:
info = await resp.json()
self.decode_dp_size = len(info.get("internal_states", [1]))
logger.info(
f"[MiniLB] DP sizes: prefill={self.prefill_dp_size}, decode={self.decode_dp_size}"
)
def _fork_dp_requests(self, request):
p_rank = random.randint(0, self.prefill_dp_size - 1)
d_rank = random.randint(0, self.decode_dp_size - 1)
prefill_req = request.copy()
decode_req = request.copy()
prefill_req["routed_dp_rank"] = p_rank
decode_req["routed_dp_rank"] = d_rank
decode_req["disagg_prefill_dp_rank"] = p_rank
return prefill_req, decode_req, d_rank
def select_pair(self):
assert len(self.prefill_urls) > 0, "No prefill servers available"
assert len(self.decode_urls) > 0, "No decode servers available"
@@ -111,6 +142,16 @@ class MiniLoadBalancer:
) -> ORJSONResponse:
assert endpoint[0] != "/", f"Endpoint should not start with '/': {endpoint}"
expected_decode_dp_rank = None
if self.test_external_dp_routing:
await self._ensure_dp_sizes()
prefill_req, decode_req, expected_decode_dp_rank = self._fork_dp_requests(
modified_request
)
else:
prefill_req = modified_request
decode_req = modified_request
async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(
total=self.timeout
@@ -130,12 +171,12 @@ class MiniLoadBalancer:
tasks = [
session.post(
f"{prefill_server}/{endpoint}",
json=modified_request,
json=prefill_req,
headers=headers,
),
session.post(
f"{decode_server}/{endpoint}",
json=modified_request,
json=decode_req,
headers=headers,
),
]
@@ -169,6 +210,16 @@ class MiniLoadBalancer:
)
trace_req_finish(bootstrap_room)
if expected_decode_dp_rank is not None:
actual = ret_json.get("meta_info", {}).get("dp_rank")
if actual != expected_decode_dp_rank:
return ORJSONResponse(
content={
"error": f"DP rank mismatch: expected {expected_decode_dp_rank}, got {actual}"
},
status_code=500,
)
return ORJSONResponse(
content=ret_json,
status_code=decode_response.status,
@@ -177,6 +228,10 @@ class MiniLoadBalancer:
async def generate_stream(
self, modified_request, prefill_server, decode_server, endpoint="generate"
):
if self.test_external_dp_routing:
warnings.warn("--test-external-dp-routing is not supported with streaming")
assert endpoint[0] != "/", f"Endpoint should not start with '/': {endpoint}"
async def stream_results():
@@ -18,6 +18,7 @@ class RouterArgs:
# PD-specific configuration
mini_lb: bool = False
test_external_dp_routing: bool = False
pd_disaggregation: bool = False # Enable PD disaggregated mode
prefill_urls: List[tuple] = dataclasses.field(
default_factory=list
@@ -360,6 +361,11 @@ class RouterArgs:
action="store_true",
help="Enable MiniLB",
)
pd_group.add_argument(
f"--{prefix}test-external-dp-routing",
action="store_true",
help="(MiniLB only) Randomly assign routed_dp_rank / disagg_prefill_dp_rank per request and verify the response dp_rank matches.",
)
pd_group.add_argument(
f"--{prefix}pd-disaggregation",
action="store_true",