fix(dp_ctrl): fix imbalanced DP rank assignment

1. feat:update total tokens immediately

fea:record request in da controller local
This commit is contained in:
laoyao0822
2026-03-26 00:12:46 +08:00
committed by wxiwnd
parent cc11cac77c
commit d9b6e90b35
5 changed files with 210 additions and 20 deletions

View File

@@ -36,6 +36,7 @@ from sglang.srt.managers.io_struct import (
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
WatchLoadUpdateReq,
DpRequestInfoqOutput,
)
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import run_scheduler_process
@@ -57,6 +58,9 @@ from sglang.srt.utils.network import NetworkAddress, bind_port, get_zmq_socket
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.srt.utils.watchdog import Watchdog
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
from sglang.srt.disaggregation.utils import (
DisaggregationMode
)
logger = logging.getLogger(__name__)
@@ -77,35 +81,137 @@ class LoadBalanceMethod(Enum):
except KeyError as exc:
raise ValueError(f"Invalid load balance method: {method}") from exc
import random
class DPBudget:
def __init__(self, dp_size: int):
def __init__(self, server_args: ServerArgs):
dp_size = server_args.dp_size
self.dp_size = dp_size
self.total_requests = [0] * dp_size
self.total_tokens = [0] * dp_size
self.ts_tic = [0.0] * dp_size
self.pool_tokens = [0] * dp_size
self.req_sum_tokens = [0] * dp_size
self.consume_token_num_per_rank = 4000
self.disaggregation_mode = DisaggregationMode(
server_args.disaggregation_mode
)
# 每个 rank 一个账本,按 timestamp 排序
self.reqs_info = [list() for _ in range(dp_size)]
def _sort_rank_ledger(self, rank: int):
self.reqs_info[rank].sort(key=lambda x: x.timestamp)
def _append_local_req(self, rank: int, req: Req):
if req is None:
return
rid = getattr(req, "rid", None)
if rid is None:
return
num_tokens = len(req.input_ids) if getattr(req, "input_ids", None) is not None else 0
completion_tokens = (
len(req.output_ids) if getattr(req, "output_ids", None) is not None else 0
)
cached_tokens = getattr(req, "cached_tokens", 0)
info = DpRequestInfoqOutput(
rid=rid,
timestamp=time.perf_counter(),
dp_rank=rank,
num_tokens=num_tokens,
cached_tokens=cached_tokens,
completion_tokens=completion_tokens,
)
self.reqs_info[rank].append(info)
def _merge_rank_reqs(self, rank: int, incoming_reqs: List[DpRequestInfoqOutput]):
if incoming_reqs is None:
return
if len(incoming_reqs) == 0:
self.reqs_info[rank] = []
return
incoming_reqs = sorted(incoming_reqs, key=lambda x: x.timestamp)
earliest_ts = incoming_reqs[0].timestamp
# 删除本地账本中比 incoming_reqs 中最早 timestamp 更旧的请求
local_reqs = [x for x in self.reqs_info[rank] if x.timestamp >= earliest_ts]
# 按 rid 合并
merged = {x.rid: x for x in local_reqs}
for x in incoming_reqs:
merged[x.rid] = x
self.reqs_info[rank] = sorted(merged.values(), key=lambda x: x.timestamp)
def update_budget(self, load_update: WatchLoadUpdateReq):
"""Update the budget."""
for load in load_update.loads:
self.total_requests[load.dp_rank] = load.num_reqs
self.total_tokens[load.dp_rank] = load.num_tokens
if load is None:
continue
def dispatch(self, method: LoadBalanceMethod):
rank = load.dp_rank
if rank is None or rank >= self.dp_size:
continue
self.total_requests[rank] = load.num_reqs
self.total_tokens[rank] = load.num_tokens
self.ts_tic[rank] = load.ts_tic
self.pool_tokens[rank] = load.pool_tokens
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self._merge_rank_reqs(rank, load.reqs_list)
else:
self._merge_rank_reqs(rank, load.reqs_list)
self.req_sum_tokens[rank] = sum(x.num_tokens for x in self.reqs_info[rank])+self.pool_tokens[rank]
def dispatch(self, method: LoadBalanceMethod, req: Req = None):
if method == LoadBalanceMethod.TOTAL_REQUESTS:
target_rank = self.total_requests.index(min(self.total_requests))
min_req = min(self.total_requests)
candidates = [i for i, x in enumerate(self.total_requests) if x == min_req]
target_rank = random.choice(candidates)
elif method == LoadBalanceMethod.TOTAL_TOKENS:
# Use total_requests as a tie-breaker when total_tokens are equal
target_rank = min(
range(self.dp_size),
key=lambda i: (self.total_tokens[i], self.total_requests[i]),
if self.disaggregation_mode == DisaggregationMode.PREFILL:
pairs = [
(self.total_tokens[i], self.total_requests[i])
for i in range(self.dp_size)
]
min_pair = min(pairs)
candidates = [
i for i in range(self.dp_size)
if (self.total_tokens[i], self.total_requests[i]) == min_pair
]
target_rank = random.choice(candidates)
min_tokens = self.total_tokens[target_rank]
else:
pairs = [
(self.total_tokens[i], self.total_requests[i])
for i in range(self.dp_size)
]
min_pair = min(pairs)
candidates = [
i for i in range(self.dp_size)
if (self.total_tokens[i], self.total_requests[i]) == min_pair
]
target_rank = random.choice(candidates)
min_tokens = self.total_tokens[target_rank]
logger.info(
f"Dispatching to DP rank {target_rank} with rough_num_tokens={min_tokens} and req_sum_tokens={self.req_sum_tokens[target_rank]}"
)
else:
return None
# Increment the load of that worker by one as a heuristic
self.total_requests[target_rank] += 1
return target_rank
if req is not None:
self.total_tokens[target_rank] += len(req.input_ids)
self.req_sum_tokens[target_rank] += len(req.input_ids)
self._append_local_req(target_rank, req)
return target_rank
class DataParallelController:
"""A controller that dispatches requests to multiple data parallel workers."""
@@ -145,7 +251,7 @@ class DataParallelController:
self.dispatching = dispatch_lookup[self.load_balance_method]
# Load balance budget
self.dp_budget = DPBudget(server_args.dp_size)
self.dp_budget = DPBudget(server_args)
# To protect changing env vars to set CUDA_VISIBLE_DEVICES.
self.env_lock = threading.Lock()
@@ -557,7 +663,7 @@ class DataParallelController:
def total_tokens_scheduler(self, req: Req):
if self.maybe_external_dp_rank_routing(req):
return
target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_TOKENS)
target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_TOKENS,req=req)
self.workers[target_worker].send_pyobj(req)
def event_loop(self):