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:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user