Tiny support sticky routing algorithm in schedule simulator (#16355)
This commit is contained in:
@@ -15,6 +15,7 @@ from sglang.srt.debug_utils.schedule_simulator.routers import (
|
|||||||
RandomRouter,
|
RandomRouter,
|
||||||
RoundRobinRouter,
|
RoundRobinRouter,
|
||||||
RouterPolicy,
|
RouterPolicy,
|
||||||
|
StickyRouter,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.schedule_simulator.schedulers import (
|
from sglang.srt.debug_utils.schedule_simulator.schedulers import (
|
||||||
FIFOScheduler,
|
FIFOScheduler,
|
||||||
@@ -34,6 +35,7 @@ __all__ = [
|
|||||||
"RouterPolicy",
|
"RouterPolicy",
|
||||||
"RandomRouter",
|
"RandomRouter",
|
||||||
"RoundRobinRouter",
|
"RoundRobinRouter",
|
||||||
|
"StickyRouter",
|
||||||
"SchedulerPolicy",
|
"SchedulerPolicy",
|
||||||
"FIFOScheduler",
|
"FIFOScheduler",
|
||||||
"MetricRecorder",
|
"MetricRecorder",
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
|
|||||||
from sglang.srt.debug_utils.schedule_simulator.routers import (
|
from sglang.srt.debug_utils.schedule_simulator.routers import (
|
||||||
RandomRouter,
|
RandomRouter,
|
||||||
RoundRobinRouter,
|
RoundRobinRouter,
|
||||||
|
StickyRouter,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.schedule_simulator.schedulers import FIFOScheduler
|
from sglang.srt.debug_utils.schedule_simulator.schedulers import FIFOScheduler
|
||||||
from sglang.srt.debug_utils.schedule_simulator.simulator import Simulator
|
from sglang.srt.debug_utils.schedule_simulator.simulator import Simulator
|
||||||
@@ -62,7 +63,10 @@ def create_arg_parser() -> argparse.ArgumentParser:
|
|||||||
|
|
||||||
parser.add_argument("--num-gpus", type=int, default=8)
|
parser.add_argument("--num-gpus", type=int, default=8)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--router", type=str, choices=["random", "round_robin"], default="round_robin"
|
"--router",
|
||||||
|
type=str,
|
||||||
|
choices=["random", "round_robin", "sticky"],
|
||||||
|
default="round_robin",
|
||||||
)
|
)
|
||||||
parser.add_argument("--scheduler", type=str, choices=["fifo"], default="fifo")
|
parser.add_argument("--scheduler", type=str, choices=["fifo"], default="fifo")
|
||||||
parser.add_argument("--max-total-tokens", type=int, default=100000)
|
parser.add_argument("--max-total-tokens", type=int, default=100000)
|
||||||
@@ -97,11 +101,13 @@ def _load_requests(args: argparse.Namespace) -> List[SimRequest]:
|
|||||||
return requests
|
return requests
|
||||||
|
|
||||||
|
|
||||||
def _create_router(name: str):
|
def _create_router(name: str, num_gpus: int):
|
||||||
if name == "random":
|
if name == "random":
|
||||||
return RandomRouter()
|
return RandomRouter()
|
||||||
if name == "round_robin":
|
if name == "round_robin":
|
||||||
return RoundRobinRouter()
|
return RoundRobinRouter()
|
||||||
|
if name == "sticky":
|
||||||
|
return StickyRouter(num_gpus)
|
||||||
raise ValueError(f"Unknown router: {name}")
|
raise ValueError(f"Unknown router: {name}")
|
||||||
|
|
||||||
|
|
||||||
@@ -113,7 +119,7 @@ def _create_scheduler(name: str):
|
|||||||
|
|
||||||
def main(args: argparse.Namespace) -> pl.DataFrame:
|
def main(args: argparse.Namespace) -> pl.DataFrame:
|
||||||
requests = _load_requests(args)
|
requests = _load_requests(args)
|
||||||
router = _create_router(args.router)
|
router = _create_router(args.router, args.num_gpus)
|
||||||
scheduler = _create_scheduler(args.scheduler)
|
scheduler = _create_scheduler(args.scheduler)
|
||||||
|
|
||||||
sim = Simulator(
|
sim = Simulator(
|
||||||
|
|||||||
@@ -3,5 +3,6 @@ from sglang.srt.debug_utils.schedule_simulator.routers.random_router import Rand
|
|||||||
from sglang.srt.debug_utils.schedule_simulator.routers.round_robin_router import (
|
from sglang.srt.debug_utils.schedule_simulator.routers.round_robin_router import (
|
||||||
RoundRobinRouter,
|
RoundRobinRouter,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.debug_utils.schedule_simulator.routers.sticky_router import StickyRouter
|
||||||
|
|
||||||
__all__ = ["RouterPolicy", "RandomRouter", "RoundRobinRouter"]
|
__all__ = ["RouterPolicy", "RandomRouter", "RoundRobinRouter", "StickyRouter"]
|
||||||
|
|||||||
@@ -0,0 +1,26 @@
|
|||||||
|
import random
|
||||||
|
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.routers.base import RouterPolicy
|
||||||
|
|
||||||
|
|
||||||
|
class StickyRouter(RouterPolicy):
|
||||||
|
def __init__(self, num_gpus: int):
|
||||||
|
self._num_gpus = num_gpus
|
||||||
|
self._group_to_gpu = defaultdict(self._assign_gpu)
|
||||||
|
|
||||||
|
def _assign_gpu(self) -> int:
|
||||||
|
return random.randint(0, self._num_gpus - 1)
|
||||||
|
|
||||||
|
def route(
|
||||||
|
self,
|
||||||
|
incoming_request: SimRequest,
|
||||||
|
gpu_states: List[GPUState],
|
||||||
|
) -> int:
|
||||||
|
group_id = incoming_request.group_id
|
||||||
|
if group_id is None:
|
||||||
|
return random.randint(0, len(gpu_states) - 1)
|
||||||
|
return self._group_to_gpu[group_id]
|
||||||
@@ -15,6 +15,7 @@ from sglang.srt.debug_utils.schedule_simulator import (
|
|||||||
SimulationResult,
|
SimulationResult,
|
||||||
Simulator,
|
Simulator,
|
||||||
StepRecord,
|
StepRecord,
|
||||||
|
StickyRouter,
|
||||||
generate_gsp_requests,
|
generate_gsp_requests,
|
||||||
generate_random_requests,
|
generate_random_requests,
|
||||||
load_from_request_logger,
|
load_from_request_logger,
|
||||||
@@ -154,6 +155,42 @@ class TestRouters(CustomTestCase):
|
|||||||
results = [router.route(req, gpu_states) for _ in range(100)]
|
results = [router.route(req, gpu_states) 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):
|
||||||
|
router = StickyRouter(num_gpus=4)
|
||||||
|
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
|
||||||
|
reqs = [
|
||||||
|
SimRequest(request_id=f"r{i}", input_len=100, output_len=50, group_id="g0")
|
||||||
|
for i in range(10)
|
||||||
|
]
|
||||||
|
results = [router.route(req, gpu_states) for req in reqs]
|
||||||
|
self.assertEqual(len(set(results)), 1)
|
||||||
|
|
||||||
|
def test_sticky_router_no_group_fallback(self):
|
||||||
|
router = StickyRouter(num_gpus=4)
|
||||||
|
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(4)]
|
||||||
|
reqs = [
|
||||||
|
SimRequest(request_id=f"r{i}", input_len=100, output_len=50)
|
||||||
|
for i in range(100)
|
||||||
|
]
|
||||||
|
results = [router.route(req, gpu_states) for req in reqs]
|
||||||
|
self.assertTrue(all(0 <= r < 4 for r in results))
|
||||||
|
|
||||||
|
def test_sticky_router_multiple_groups(self):
|
||||||
|
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"]:
|
||||||
|
reqs = [
|
||||||
|
SimRequest(
|
||||||
|
request_id=f"{group_id}_r{i}",
|
||||||
|
input_len=100,
|
||||||
|
output_len=50,
|
||||||
|
group_id=group_id,
|
||||||
|
)
|
||||||
|
for i in range(5)
|
||||||
|
]
|
||||||
|
results = [router.route(req, gpu_states) for req in reqs]
|
||||||
|
self.assertEqual(len(set(results)), 1)
|
||||||
|
|
||||||
|
|
||||||
class TestFIFOScheduler(CustomTestCase):
|
class TestFIFOScheduler(CustomTestCase):
|
||||||
def test_runs_pending_requests(self):
|
def test_runs_pending_requests(self):
|
||||||
@@ -508,6 +545,54 @@ class TestCLI(CustomTestCase):
|
|||||||
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
|
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
|
||||||
self.assertIn("router=random", result.stdout)
|
self.assertIn("router=random", result.stdout)
|
||||||
|
|
||||||
|
def test_cli_sticky_router(self):
|
||||||
|
result = self._run_cli(
|
||||||
|
"--synth-gsp",
|
||||||
|
"--synth-gsp-num-groups",
|
||||||
|
"2",
|
||||||
|
"--synth-gsp-prompts-per-group",
|
||||||
|
"3",
|
||||||
|
"--synth-seed",
|
||||||
|
"42",
|
||||||
|
"--num-gpus",
|
||||||
|
"4",
|
||||||
|
"--router",
|
||||||
|
"sticky",
|
||||||
|
)
|
||||||
|
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
|
||||||
|
self.assertIn("router=sticky", result.stdout)
|
||||||
|
|
||||||
|
def test_e2e_sticky_router_group_locality(self):
|
||||||
|
result = self._run_cli(
|
||||||
|
"--synth-gsp",
|
||||||
|
"--synth-gsp-num-groups",
|
||||||
|
"2",
|
||||||
|
"--synth-gsp-prompts-per-group",
|
||||||
|
"2",
|
||||||
|
"--synth-gsp-system-prompt-len",
|
||||||
|
"10",
|
||||||
|
"--synth-gsp-question-len",
|
||||||
|
"10",
|
||||||
|
"--synth-gsp-output-len",
|
||||||
|
"2",
|
||||||
|
"--synth-seed",
|
||||||
|
"123",
|
||||||
|
"--num-gpus",
|
||||||
|
"2",
|
||||||
|
"--router",
|
||||||
|
"sticky",
|
||||||
|
"--max-total-tokens",
|
||||||
|
"1000",
|
||||||
|
"--log-level",
|
||||||
|
"2",
|
||||||
|
)
|
||||||
|
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
|
||||||
|
for expected in [
|
||||||
|
"step=0 | GPU0[R=2:gsp0,gsp1 Q=0:-] | GPU1[R=2:gsp2,gsp3 Q=0:-]",
|
||||||
|
"step=1 | GPU0[R=0:- Q=0:-] | GPU1[R=0:- Q=0:-]",
|
||||||
|
]:
|
||||||
|
self.assertIn(expected, result.stdout)
|
||||||
|
|
||||||
def test_cli_synthetic(self):
|
def test_cli_synthetic(self):
|
||||||
result = self._run_cli(
|
result = self._run_cli(
|
||||||
"--synthetic",
|
"--synthetic",
|
||||||
|
|||||||
Reference in New Issue
Block a user