Support multiple engines for router simulation in schedule simulator (#16368)

This commit is contained in:
fzyzcjy
2026-01-04 11:46:05 +08:00
committed by GitHub
parent e797f0c570
commit d7aa0ce72f
7 changed files with 72 additions and 93 deletions
@@ -63,7 +63,8 @@ def create_arg_parser() -> argparse.ArgumentParser:
parser.add_argument("--synth-gsp-output-len", type=int, default=256) parser.add_argument("--synth-gsp-output-len", type=int, default=256)
parser.add_argument("--synth-gsp-range-ratio", type=float, default=1.0) parser.add_argument("--synth-gsp-range-ratio", type=float, default=1.0)
parser.add_argument("--num-gpus", type=int, default=8) parser.add_argument("--num-gpus-per-engine", type=int, default=8)
parser.add_argument("--num-engines", type=int, default=1)
parser.add_argument( parser.add_argument(
"--router", "--router",
type=str, type=str,
@@ -111,13 +112,13 @@ def _load_requests(args: argparse.Namespace) -> List[SimRequest]:
return requests return requests
def _create_router(name: str, num_gpus: int): def _create_router(name: str, total_gpus: int):
if name == "random": if name == "random":
return RandomRouter() return RandomRouter(total_gpus)
if name == "round_robin": if name == "round_robin":
return RoundRobinRouter() return RoundRobinRouter(total_gpus)
if name == "sticky": if name == "sticky":
return StickyRouter(num_gpus) return StickyRouter(total_gpus)
raise ValueError(f"Unknown router: {name}") raise ValueError(f"Unknown router: {name}")
@@ -131,11 +132,12 @@ def main(args: argparse.Namespace) -> SimulationResult:
if args.synth_seed is not None: if args.synth_seed is not None:
random.seed(args.synth_seed) random.seed(args.synth_seed)
requests = _load_requests(args) requests = _load_requests(args)
router = _create_router(args.router, args.num_gpus) total_gpus = args.num_gpus_per_engine * args.num_engines
router = _create_router(args.router, total_gpus)
scheduler = _create_scheduler(args.scheduler) scheduler = _create_scheduler(args.scheduler)
sim = Simulator( sim = Simulator(
num_gpus=args.num_gpus, num_gpus_per_engine=args.num_gpus_per_engine,
router=router, router=router,
scheduler=scheduler, scheduler=scheduler,
recorders=[ recorders=[
@@ -150,7 +152,7 @@ def main(args: argparse.Namespace) -> SimulationResult:
) )
print( print(
f"Running simulation with {args.num_gpus} GPUs, router={args.router}, scheduler={args.scheduler}" f"Running simulation with {args.num_gpus_per_engine} GPUs/engine x {args.num_engines} engines, router={args.router}, scheduler={args.scheduler}"
) )
result = sim.run(requests) result = sim.run(requests)
@@ -1,14 +1,8 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import List
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
class RouterPolicy(ABC): class RouterPolicy(ABC):
@abstractmethod @abstractmethod
def route( def route(self, incoming_request: SimRequest) -> int: ...
self,
incoming_request: SimRequest,
gpu_states: List[GPUState],
) -> int: ...
@@ -1,15 +1,12 @@
import random import random
from typing import List
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy
class RandomRouter(RouterPolicy): class RandomRouter(RouterPolicy):
def route( def __init__(self, num_gpus: int):
self, self._num_gpus = num_gpus
incoming_request: SimRequest,
gpu_states: List[GPUState], def route(self, incoming_request: SimRequest) -> int:
) -> int: return random.randint(0, self._num_gpus - 1)
return random.randint(0, len(gpu_states) - 1)
@@ -1,19 +1,13 @@
from typing import List
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy
class RoundRobinRouter(RouterPolicy): class RoundRobinRouter(RouterPolicy):
def __init__(self): def __init__(self, num_gpus: int):
self._num_gpus = num_gpus
self._counter = 0 self._counter = 0
def route( def route(self, incoming_request: SimRequest) -> int:
self, gpu_id = self._counter % self._num_gpus
incoming_request: SimRequest,
gpu_states: List[GPUState],
) -> int:
gpu_id = self._counter % len(gpu_states)
self._counter += 1 self._counter += 1
return gpu_id return gpu_id
@@ -1,8 +1,6 @@
import random import random
from collections import defaultdict from collections import defaultdict
from typing import List
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy
@@ -15,12 +13,8 @@ class StickyRouter(RouterPolicy):
def _assign_gpu(self) -> int: def _assign_gpu(self) -> int:
return random.randint(0, self._num_gpus - 1) return random.randint(0, self._num_gpus - 1)
def route( def route(self, incoming_request: SimRequest) -> int:
self,
incoming_request: SimRequest,
gpu_states: List[GPUState],
) -> int:
group_id = incoming_request.group_id group_id = incoming_request.group_id
if group_id is None: if group_id is None:
return random.randint(0, len(gpu_states) - 1) return random.randint(0, self._num_gpus - 1)
return self._group_to_gpu[group_id] return self._group_to_gpu[group_id]
@@ -17,7 +17,7 @@ class SimulationResult:
class Simulator: class Simulator:
def __init__( def __init__(
self, self,
num_gpus: int, num_gpus_per_engine: int,
router: RouterPolicy, router: RouterPolicy,
scheduler: SchedulerPolicy, scheduler: SchedulerPolicy,
recorders: Optional[List[MetricRecorder]] = None, recorders: Optional[List[MetricRecorder]] = None,
@@ -26,7 +26,7 @@ class Simulator:
stop_criteria: str = "all_done", stop_criteria: str = "all_done",
max_steps: Optional[int] = None, max_steps: Optional[int] = None,
): ):
self.num_gpus = num_gpus self.num_gpus_per_engine = num_gpus_per_engine
self.router = router self.router = router
self.scheduler = scheduler self.scheduler = scheduler
self.recorders = recorders or [] self.recorders = recorders or []
@@ -40,7 +40,7 @@ class Simulator:
def run(self, requests: List[SimRequest]) -> SimulationResult: def run(self, requests: List[SimRequest]) -> SimulationResult:
self.gpu_states = [ self.gpu_states = [
GPUState(gpu_id=i, max_total_tokens=self.max_total_tokens) GPUState(gpu_id=i, max_total_tokens=self.max_total_tokens)
for i in range(self.num_gpus) for i in range(self.num_gpus_per_engine)
] ]
self.step = 0 self.step = 0
step_records: List[StepRecord] = [] step_records: List[StepRecord] = []
@@ -75,7 +75,8 @@ class Simulator:
def _route_requests(self, incoming_requests: List[SimRequest]) -> None: def _route_requests(self, incoming_requests: List[SimRequest]) -> None:
for req in incoming_requests: for req in incoming_requests:
gpu_id = self.router.route(req, self.gpu_states) gpu_id = self.router.route(req)
if gpu_id < self.num_gpus_per_engine:
self.gpu_states[gpu_id].pending_requests.append(req) self.gpu_states[gpu_id].pending_requests.append(req)
def _schedule_all_gpus(self) -> None: def _schedule_all_gpus(self) -> None:
@@ -144,42 +144,37 @@ class TestGPUState(CustomTestCase):
class TestRouters(CustomTestCase): class TestRouters(CustomTestCase):
def test_round_robin(self): def test_round_robin(self):
router = RoundRobinRouter() router = RoundRobinRouter(num_gpus=4)
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
req = SimRequest(request_id="r1", input_len=100, output_len=50) req = SimRequest(request_id="r1", input_len=100, output_len=50)
results = [router.route(req, gpu_states) for _ in range(8)] results = [router.route(req) for _ in range(8)]
self.assertEqual(results, [0, 1, 2, 3, 0, 1, 2, 3]) self.assertEqual(results, [0, 1, 2, 3, 0, 1, 2, 3])
def test_random_router(self): def test_random_router(self):
router = RandomRouter() router = RandomRouter(num_gpus=4)
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
req = SimRequest(request_id="r1", input_len=100, output_len=50) req = SimRequest(request_id="r1", input_len=100, output_len=50)
results = [router.route(req, gpu_states) for _ in range(100)] results = [router.route(req) for _ in range(100)]
self.assertTrue(all(0 <= r < 4 for r in results)) self.assertTrue(all(0 <= r < 4 for r in results))
def test_sticky_router_same_group_same_gpu(self): def test_sticky_router_same_group_same_gpu(self):
router = StickyRouter(num_gpus=4) router = StickyRouter(num_gpus=4)
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
reqs = [ reqs = [
SimRequest(request_id=f"r{i}", input_len=100, output_len=50, group_id="g0") SimRequest(request_id=f"r{i}", input_len=100, output_len=50, group_id="g0")
for i in range(10) for i in range(10)
] ]
results = [router.route(req, gpu_states) for req in reqs] results = [router.route(req) for req in reqs]
self.assertEqual(len(set(results)), 1) self.assertEqual(len(set(results)), 1)
def test_sticky_router_no_group_fallback(self): def test_sticky_router_no_group_fallback(self):
router = StickyRouter(num_gpus=4) router = StickyRouter(num_gpus=4)
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
reqs = [ reqs = [
SimRequest(request_id=f"r{i}", input_len=100, output_len=50) SimRequest(request_id=f"r{i}", input_len=100, output_len=50)
for i in range(100) for i in range(100)
] ]
results = [router.route(req, gpu_states) for req in reqs] results = [router.route(req) for req in reqs]
self.assertTrue(all(0 <= r < 4 for r in results)) self.assertTrue(all(0 <= r < 4 for r in results))
def test_sticky_router_multiple_groups(self): def test_sticky_router_multiple_groups(self):
router = StickyRouter(num_gpus=4) router = StickyRouter(num_gpus=4)
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
for group_id in ["g0", "g1", "g2"]: for group_id in ["g0", "g1", "g2"]:
reqs = [ reqs = [
SimRequest( SimRequest(
@@ -190,7 +185,7 @@ class TestRouters(CustomTestCase):
) )
for i in range(5) for i in range(5)
] ]
results = [router.route(req, gpu_states) for req in reqs] results = [router.route(req) for req in reqs]
self.assertEqual(len(set(results)), 1) self.assertEqual(len(set(results)), 1)
@@ -405,8 +400,8 @@ class TestSimulator(CustomTestCase):
for i in range(10) for i in range(10)
] ]
sim = Simulator( sim = Simulator(
num_gpus=2, num_gpus_per_engine=2,
router=RoundRobinRouter(), router=RoundRobinRouter(num_gpus=2),
scheduler=FIFOScheduler(), scheduler=FIFOScheduler(),
recorders=[ recorders=[
BatchSizeBalancednessRecorder(), BatchSizeBalancednessRecorder(),
@@ -424,8 +419,8 @@ class TestSimulator(CustomTestCase):
SimRequest(request_id=f"r{i}", input_len=10, output_len=3) for i in range(4) SimRequest(request_id=f"r{i}", input_len=10, output_len=3) for i in range(4)
] ]
sim = Simulator( sim = Simulator(
num_gpus=2, num_gpus_per_engine=2,
router=RoundRobinRouter(), router=RoundRobinRouter(num_gpus=2),
scheduler=FIFOScheduler(), scheduler=FIFOScheduler(),
max_total_tokens=10000, max_total_tokens=10000,
) )
@@ -436,7 +431,9 @@ class TestSimulator(CustomTestCase):
def test_empty_requests(self): def test_empty_requests(self):
sim = Simulator( sim = Simulator(
num_gpus=2, router=RoundRobinRouter(), scheduler=FIFOScheduler() num_gpus_per_engine=2,
router=RoundRobinRouter(num_gpus=2),
scheduler=FIFOScheduler(),
) )
result = sim.run([]) result = sim.run([])
self.assertEqual(result.summary, {}) self.assertEqual(result.summary, {})
@@ -447,8 +444,8 @@ class TestSimulator(CustomTestCase):
SimRequest(request_id=f"r{i}", input_len=10, output_len=3) for i in range(4) SimRequest(request_id=f"r{i}", input_len=10, output_len=3) for i in range(4)
] ]
sim = Simulator( sim = Simulator(
num_gpus=2, num_gpus_per_engine=2,
router=RoundRobinRouter(), router=RoundRobinRouter(num_gpus=2),
scheduler=FIFOScheduler(), scheduler=FIFOScheduler(),
max_total_tokens=10000, max_total_tokens=10000,
) )
@@ -460,23 +457,18 @@ class TestSimulator(CustomTestCase):
self.assertEqual(len([r for r in result.step_records if r.step == 0]), 2) self.assertEqual(len([r for r in result.step_records if r.step == 0]), 2)
def test_preemption_due_to_token_growth(self): def test_preemption_due_to_token_growth(self):
# 2 requests on 1 GPU, each input_len=50, output_len=10
# max_total_tokens=110, so initially both can run (100 tokens)
# After 5 decode steps, total = 50+5 + 50+5 = 110, still ok
# After 6 decode steps, total = 50+6 + 50+6 = 112 > 110, need preempt
requests = [ requests = [
SimRequest(request_id="r0", input_len=50, output_len=10), SimRequest(request_id="r0", input_len=50, output_len=10),
SimRequest(request_id="r1", input_len=50, output_len=10), SimRequest(request_id="r1", input_len=50, output_len=10),
] ]
sim = Simulator( sim = Simulator(
num_gpus=1, num_gpus_per_engine=1,
router=RoundRobinRouter(), router=RoundRobinRouter(num_gpus=1),
scheduler=FIFOScheduler(), scheduler=FIFOScheduler(),
max_total_tokens=110, max_total_tokens=110,
) )
result = sim.run(requests) result = sim.run(requests)
# Check that preemption happened at some point
found_preemption = False found_preemption = False
for record in result.step_records: for record in result.step_records:
if record.running_count == 1 and record.pending_count == 1: if record.running_count == 1 and record.pending_count == 1:
@@ -523,7 +515,7 @@ class TestCLI(CustomTestCase):
output_file = f.name output_file = f.name
result = self._run_cli( result = self._run_cli(
"--input", input_file, "--num-gpus", "2", "--output", output_file "--input", input_file, "--num-gpus-per-engine", "2", "--output", output_file
) )
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}") self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self.assertIn("Loaded 2 requests", result.stdout) self.assertIn("Loaded 2 requests", result.stdout)
@@ -562,7 +554,7 @@ class TestCLI(CustomTestCase):
"2", "2",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"2", "2",
"--router", "--router",
"sticky", "sticky",
@@ -586,7 +578,7 @@ class TestCLI(CustomTestCase):
"128", "128",
"--synth-random-range-ratio", "--synth-random-range-ratio",
"0.5", "0.5",
"--num-gpus", "--num-gpus-per-engine",
"4", "4",
) )
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}") self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
@@ -599,7 +591,7 @@ class TestCLI(CustomTestCase):
"10", "10",
"--synth-random-output-len", "--synth-random-output-len",
"5", "5",
"--num-gpus", "--num-gpus-per-engine",
"2", "2",
"--log-level", "--log-level",
"1", "1",
@@ -608,7 +600,6 @@ class TestCLI(CustomTestCase):
self.assertIn("step=", result.stdout) self.assertIn("step=", result.stdout)
def test_e2e_simple_no_queuing(self): def test_e2e_simple_no_queuing(self):
# 4 requests, input_len=10, output_len=2, 2 GPUs, all fit in memory
result = self._run_cli( result = self._run_cli(
"--synthetic", "--synthetic",
"--synth-random-num-requests", "--synth-random-num-requests",
@@ -621,7 +612,7 @@ class TestCLI(CustomTestCase):
"1.0", "1.0",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"2", "2",
"--max-total-tokens", "--max-total-tokens",
"10000", "10000",
@@ -651,7 +642,7 @@ class TestCLI(CustomTestCase):
"1.0", "1.0",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"1", "1",
"--max-total-tokens", "--max-total-tokens",
"210", "210",
@@ -683,7 +674,7 @@ step=5 | GPU0[R=0:- Q=0:-]""",
"1.0", "1.0",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"1", "1",
"--max-total-tokens", "--max-total-tokens",
"110", "110",
@@ -717,7 +708,7 @@ step=13 | GPU0[R=0:- Q=0:-]""",
"10", "10",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"2", "2",
) )
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}") self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
@@ -741,7 +732,7 @@ step=13 | GPU0[R=0:- Q=0:-]""",
"2", "2",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"1", "1",
"--max-total-tokens", "--max-total-tokens",
"80", "80",
@@ -769,21 +760,25 @@ class TestLargerScale(CustomTestCase):
result = self._run_main( result = self._run_main(
"--synthetic", "--synthetic",
"--synth-random-num-requests", "--synth-random-num-requests",
"2000", "500000",
"--synth-random-input-len", "--synth-random-input-len",
"32000", "32000",
"--synth-random-output-len", "--synth-random-output-len",
"2000", "2000",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"8", "8",
"--num-engines",
"250",
"--router", "--router",
"random", "random",
"--max-total-tokens", "--max-total-tokens",
"2000000", "2000000",
"--stop-criteria", "--stop-criteria",
"exist_no_pending", "exist_no_pending",
"--max-steps",
"1500",
) )
self._assert_in_range( self._assert_in_range(
result.summary["attention_compute_balancedness_mean"], 0.95, 1.0, "attn" result.summary["attention_compute_balancedness_mean"], 0.95, 1.0, "attn"
@@ -797,9 +792,9 @@ class TestLargerScale(CustomTestCase):
return self._run_main( return self._run_main(
"--synth-gsp", "--synth-gsp",
"--synth-gsp-num-groups", "--synth-gsp-num-groups",
"200", "50000",
"--synth-gsp-prompts-per-group", "--synth-gsp-prompts-per-group",
"20", "100",
"--synth-gsp-system-prompt-len", "--synth-gsp-system-prompt-len",
"31000", "31000",
"--synth-gsp-question-len", "--synth-gsp-question-len",
@@ -808,8 +803,10 @@ class TestLargerScale(CustomTestCase):
"8000", "8000",
"--synth-seed", "--synth-seed",
"42", "42",
"--num-gpus", "--num-gpus-per-engine",
"8", "8",
"--num-engines",
"250",
"--router", "--router",
router, router,
"--max-total-tokens", "--max-total-tokens",
@@ -823,22 +820,22 @@ class TestLargerScale(CustomTestCase):
def test_gsp_workload_random_policy(self): def test_gsp_workload_random_policy(self):
result = self._run_gsp_workload("random") result = self._run_gsp_workload("random")
self._assert_in_range( self._assert_in_range(
result.summary["attention_compute_balancedness_mean"], 0.88, 0.98, "attn" result.summary["attention_compute_balancedness_mean"], 0.90, 0.97, "attn"
) )
self._assert_in_range( self._assert_in_range(
result.summary["batch_size_balancedness_mean"], 0.88, 0.98, "bs" result.summary["batch_size_balancedness_mean"], 0.90, 0.97, "bs"
) )
self._assert_in_range(result.summary["avg_batch_size"], 24, 27, "avg_bs") self._assert_in_range(result.summary["avg_batch_size"], 14, 17, "avg_bs")
def test_gsp_workload_sticky_policy(self): def test_gsp_workload_sticky_policy(self):
result = self._run_gsp_workload("sticky") result = self._run_gsp_workload("sticky")
self._assert_in_range( self._assert_in_range(
result.summary["attention_compute_balancedness_mean"], 0.63, 0.70, "attn" result.summary["attention_compute_balancedness_mean"], 0.64, 0.71, "attn"
) )
self._assert_in_range( self._assert_in_range(
result.summary["batch_size_balancedness_mean"], 0.63, 0.71, "bs" result.summary["batch_size_balancedness_mean"], 0.64, 0.71, "bs"
) )
self._assert_in_range(result.summary["avg_batch_size"], 33, 37, "avg_bs") self._assert_in_range(result.summary["avg_batch_size"], 31, 36, "avg_bs")
if __name__ == "__main__": if __name__ == "__main__":