Support simple schedule simulator (#16344)

This commit is contained in:
fzyzcjy
2026-01-03 22:05:40 +08:00
committed by GitHub
parent b65c9889a0
commit 87ef05e2e1
18 changed files with 1064 additions and 0 deletions

View File

@@ -0,0 +1,45 @@
from sglang.srt.debug_utils.schedule_simulator.data_source import (
generate_random_requests,
load_from_request_logger,
)
from sglang.srt.debug_utils.schedule_simulator.entrypoint import create_arg_parser, main
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState, StepRecord
from sglang.srt.debug_utils.schedule_simulator.metrics import (
AttentionBalancednessRecorder,
BatchSizeBalancednessRecorder,
MetricRecorder,
)
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
from sglang.srt.debug_utils.schedule_simulator.routers import (
RandomRouter,
RoundRobinRouter,
RouterPolicy,
)
from sglang.srt.debug_utils.schedule_simulator.schedulers import (
FIFOScheduler,
SchedulerPolicy,
)
from sglang.srt.debug_utils.schedule_simulator.simulator import (
SimulationResult,
Simulator,
)
__all__ = [
"SimRequest",
"GPUState",
"Simulator",
"SimulationResult",
"StepRecord",
"RouterPolicy",
"RandomRouter",
"RoundRobinRouter",
"SchedulerPolicy",
"FIFOScheduler",
"MetricRecorder",
"BatchSizeBalancednessRecorder",
"AttentionBalancednessRecorder",
"load_from_request_logger",
"generate_random_requests",
"create_arg_parser",
"main",
]

View File

@@ -0,0 +1,8 @@
from sglang.srt.debug_utils.schedule_simulator.entrypoint import create_arg_parser, main
if __name__ == "__main__":
parser = create_arg_parser()
args = parser.parse_args()
df = main(args)
print(f"\nDataFrame shape: {df.shape}")
print(df.head(20))

View File

@@ -0,0 +1,8 @@
from sglang.srt.debug_utils.schedule_simulator.data_source.data_loader import (
load_from_request_logger,
)
from sglang.srt.debug_utils.schedule_simulator.data_source.data_synthesis import (
generate_random_requests,
)
__all__ = ["load_from_request_logger", "generate_random_requests"]

View File

@@ -0,0 +1,34 @@
import json
from pathlib import Path
from typing import List, Union
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
def load_from_request_logger(file_path: Union[str, Path]) -> List[SimRequest]:
requests = []
file_path = Path(file_path)
with file_path.open(encoding="utf-8") as f:
for line_num, line in enumerate(f):
line = line.strip()
if not line or not line.startswith("{"):
continue
data = json.loads(line)
if data.get("event") != "request.finished":
continue
rid = data.get("rid", f"req_{line_num}")
meta_info = data["out"]["meta_info"]
requests.append(
SimRequest(
request_id=rid,
input_len=meta_info["prompt_tokens"],
output_len=meta_info["completion_tokens"],
)
)
return requests

View File

@@ -0,0 +1,34 @@
import random
from typing import List, Optional
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
def generate_random_requests(
num_requests: int,
input_len: int,
output_len: int,
range_ratio: float = 1.0,
seed: Optional[int] = None,
) -> List[SimRequest]:
if seed is not None:
random.seed(seed)
requests = []
for i in range(num_requests):
isl = _random_len(input_len, range_ratio)
osl = _random_len(output_len, range_ratio)
requests.append(
SimRequest(
request_id=f"syn{i}",
input_len=isl,
output_len=osl,
)
)
return requests
def _random_len(full_len: int, range_ratio: float) -> int:
min_len = max(int(full_len * range_ratio), 1)
return random.randint(min_len, full_len)

View File

@@ -0,0 +1,123 @@
import argparse
import json
from dataclasses import asdict
from typing import List
import polars as pl
from sglang.srt.debug_utils.schedule_simulator.data_source.data_loader import (
load_from_request_logger,
)
from sglang.srt.debug_utils.schedule_simulator.data_source.data_synthesis import (
generate_random_requests,
)
from sglang.srt.debug_utils.schedule_simulator.metrics import (
AttentionBalancednessRecorder,
BatchSizeBalancednessRecorder,
)
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
from sglang.srt.debug_utils.schedule_simulator.routers import (
RandomRouter,
RoundRobinRouter,
)
from sglang.srt.debug_utils.schedule_simulator.schedulers import FIFOScheduler
from sglang.srt.debug_utils.schedule_simulator.simulator import Simulator
def create_arg_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Schedule Simulator for analyzing request scheduling across GPUs"
)
data_group = parser.add_mutually_exclusive_group(required=True)
data_group.add_argument(
"--input", type=str, help="Path to request_logger JSON file"
)
data_group.add_argument(
"--synthetic", action="store_true", help="Use synthetic data generation"
)
parser.add_argument("--synth-num-requests", type=int, default=1000)
parser.add_argument("--synth-input-len", type=int, default=1024)
parser.add_argument("--synth-output-len", type=int, default=256)
parser.add_argument("--synth-range-ratio", type=float, default=1.0)
parser.add_argument("--synth-seed", type=int, default=None)
parser.add_argument("--num-gpus", type=int, default=8)
parser.add_argument(
"--router", type=str, choices=["random", "round_robin"], default="round_robin"
)
parser.add_argument("--scheduler", type=str, choices=["fifo"], default="fifo")
parser.add_argument("--max-total-tokens", type=int, default=100000)
parser.add_argument("--output", type=str, default=None)
parser.add_argument("--log-level", type=int, choices=[0, 1, 2], default=0)
return parser
def _load_requests(args: argparse.Namespace) -> List[SimRequest]:
if args.input:
requests = load_from_request_logger(args.input)
print(f"Loaded {len(requests)} requests from {args.input}")
else:
requests = generate_random_requests(
num_requests=args.synth_num_requests,
input_len=args.synth_input_len,
output_len=args.synth_output_len,
range_ratio=args.synth_range_ratio,
seed=args.synth_seed,
)
print(
f"Generated {len(requests)} synthetic requests "
f"(synth_input_len={args.synth_input_len}, "
f"synth_output_len={args.synth_output_len}, "
f"synth_range_ratio={args.synth_range_ratio})"
)
return requests
def _create_router(name: str):
if name == "random":
return RandomRouter()
if name == "round_robin":
return RoundRobinRouter()
raise ValueError(f"Unknown router: {name}")
def _create_scheduler(name: str):
if name == "fifo":
return FIFOScheduler()
raise ValueError(f"Unknown scheduler: {name}")
def main(args: argparse.Namespace) -> pl.DataFrame:
requests = _load_requests(args)
router = _create_router(args.router)
scheduler = _create_scheduler(args.scheduler)
sim = Simulator(
num_gpus=args.num_gpus,
router=router,
scheduler=scheduler,
recorders=[BatchSizeBalancednessRecorder(), AttentionBalancednessRecorder()],
log_level=args.log_level,
max_total_tokens=args.max_total_tokens,
)
print(
f"Running simulation with {args.num_gpus} GPUs, router={args.router}, scheduler={args.scheduler}"
)
result = sim.run(requests)
df = pl.DataFrame([asdict(r) for r in result.step_records])
print("\n=== Summary ===")
for key, value in result.summary.items():
print(f"{key}: {value:.4f}" if isinstance(value, float) else f"{key}: {value}")
if args.output:
with open(args.output, "w") as f:
json.dump(result.summary, f, indent=2)
print(f"\nSummary saved to {args.output}")
return df

View File

@@ -0,0 +1,60 @@
from dataclasses import dataclass, field
from typing import List
from sglang.srt.debug_utils.schedule_simulator.request import SimRequest
@dataclass
class StepRecord:
step: int
gpu_id: int
running_count: int
pending_count: int
total_seq_len: int
running_req_ids: List[str] = field(default_factory=list)
pending_req_ids: List[str] = field(default_factory=list)
@dataclass
class GPUState:
gpu_id: int
max_total_tokens: int
pending_requests: List[SimRequest] = field(default_factory=list)
running_requests: List[SimRequest] = field(default_factory=list)
def batch_size(self) -> int:
return len(self.running_requests)
def total_seq_len(self) -> int:
return sum(req.seq_len() for req in self.running_requests)
def is_valid(self) -> bool:
return self.total_seq_len() <= self.max_total_tokens
def start_request(self, req: SimRequest) -> None:
assert req in self.pending_requests
self.pending_requests.remove(req)
self.running_requests.append(req)
def evict_request(self, req: SimRequest) -> None:
assert req in self.running_requests
self.running_requests.remove(req)
self.pending_requests.insert(0, req)
def execute_step(self) -> None:
for req in self.running_requests:
req.decoded_tokens += 1
self.running_requests = [
r for r in self.running_requests if not r.is_finished()
]
def get_step_record(self, step: int) -> StepRecord:
return StepRecord(
step=step,
gpu_id=self.gpu_id,
running_count=len(self.running_requests),
pending_count=len(self.pending_requests),
total_seq_len=self.total_seq_len(),
running_req_ids=[r.request_id for r in self.running_requests],
pending_req_ids=[r.request_id for r in self.pending_requests],
)

View File

@@ -0,0 +1,45 @@
from abc import ABC, abstractmethod
from typing import Any, Callable, Dict, List
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
class MetricRecorder(ABC):
@abstractmethod
def on_step_end(self, step: int, gpu_states: List[GPUState]) -> None: ...
@abstractmethod
def get_summary(self) -> Dict[str, Any]: ...
class BalancednessRecorder(MetricRecorder):
def __init__(self, name: str, value_fn: Callable[[GPUState], float]):
self._name = name
self._value_fn = value_fn
self._history: List[float] = []
def on_step_end(self, step: int, gpu_states: List[GPUState]) -> None:
values = [self._value_fn(gpu) for gpu in gpu_states]
max_val = max(values) if values else 0
mean_val = sum(values) / len(values) if values else 0
balancedness = mean_val / max_val if max_val > 0 else 1.0
self._history.append(balancedness)
def get_summary(self) -> Dict[str, Any]:
if not self._history:
return {f"{self._name}_mean": 0.0}
return {
f"{self._name}_mean": sum(self._history) / len(self._history),
f"{self._name}_min": min(self._history),
f"{self._name}_max": max(self._history),
}
def BatchSizeBalancednessRecorder() -> BalancednessRecorder:
return BalancednessRecorder("batch_size_balancedness", lambda gpu: gpu.batch_size())
def AttentionBalancednessRecorder() -> BalancednessRecorder:
return BalancednessRecorder(
"attention_balancedness", lambda gpu: gpu.total_seq_len()
)

View File

@@ -0,0 +1,15 @@
from dataclasses import dataclass
@dataclass
class SimRequest:
request_id: str
input_len: int
output_len: int
decoded_tokens: int = 0
def seq_len(self) -> int:
return self.input_len + self.decoded_tokens
def is_finished(self) -> bool:
return self.decoded_tokens >= self.output_len

View File

@@ -0,0 +1,7 @@
from sglang.srt.debug_utils.schedule_simulator.routers.base import RouterPolicy
from sglang.srt.debug_utils.schedule_simulator.routers.random_router import RandomRouter
from sglang.srt.debug_utils.schedule_simulator.routers.round_robin_router import (
RoundRobinRouter,
)
__all__ = ["RouterPolicy", "RandomRouter", "RoundRobinRouter"]

View File

@@ -0,0 +1,14 @@
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
class RouterPolicy(ABC):
@abstractmethod
def route(
self,
incoming_request: SimRequest,
gpu_states: List[GPUState],
) -> int: ...

View File

@@ -0,0 +1,15 @@
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.routers.base import RouterPolicy
class RandomRouter(RouterPolicy):
def route(
self,
incoming_request: SimRequest,
gpu_states: List[GPUState],
) -> int:
return random.randint(0, len(gpu_states) - 1)

View File

@@ -0,0 +1,19 @@
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 RoundRobinRouter(RouterPolicy):
def __init__(self):
self._counter = 0
def route(
self,
incoming_request: SimRequest,
gpu_states: List[GPUState],
) -> int:
gpu_id = self._counter % len(gpu_states)
self._counter += 1
return gpu_id

View File

@@ -0,0 +1,6 @@
from sglang.srt.debug_utils.schedule_simulator.schedulers.base import SchedulerPolicy
from sglang.srt.debug_utils.schedule_simulator.schedulers.fifo_scheduler import (
FIFOScheduler,
)
__all__ = ["SchedulerPolicy", "FIFOScheduler"]

View File

@@ -0,0 +1,10 @@
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
class SchedulerPolicy(ABC):
@abstractmethod
def schedule(self, gpu_state: "GPUState") -> None: ...

View File

@@ -0,0 +1,16 @@
from typing import TYPE_CHECKING
from sglang.srt.debug_utils.schedule_simulator.schedulers.base import SchedulerPolicy
if TYPE_CHECKING:
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState
class FIFOScheduler(SchedulerPolicy):
def schedule(self, gpu_state: "GPUState") -> None:
while not gpu_state.is_valid() and gpu_state.running_requests:
gpu_state.evict_request(gpu_state.running_requests[-1])
for req in list(gpu_state.pending_requests):
if gpu_state.total_seq_len() + req.seq_len() <= gpu_state.max_total_tokens:
gpu_state.start_request(req)

View File

@@ -0,0 +1,109 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from sglang.srt.debug_utils.schedule_simulator.gpu_state import GPUState, StepRecord
from sglang.srt.debug_utils.schedule_simulator.metrics import MetricRecorder
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.schedulers.base import SchedulerPolicy
@dataclass
class SimulationResult:
step_records: List[StepRecord]
summary: Dict[str, Any]
class Simulator:
def __init__(
self,
num_gpus: int,
router: RouterPolicy,
scheduler: SchedulerPolicy,
recorders: Optional[List[MetricRecorder]] = None,
log_level: int = 0,
max_total_tokens: int = 100000,
):
self.num_gpus = num_gpus
self.router = router
self.scheduler = scheduler
self.recorders = recorders or []
self.log_level = log_level
self.max_total_tokens = max_total_tokens
self.gpu_states: List[GPUState] = []
self.step = 0
def run(self, requests: List[SimRequest]) -> SimulationResult:
self.gpu_states = [
GPUState(gpu_id=i, max_total_tokens=self.max_total_tokens)
for i in range(self.num_gpus)
]
self.step = 0
step_records: List[StepRecord] = []
incoming_requests = list(requests)
while self._has_work(incoming_requests):
self._route_requests(incoming_requests)
incoming_requests.clear()
self._schedule_all_gpus()
self._execute_step()
step_records.extend(
gpu.get_step_record(self.step) for gpu in self.gpu_states
)
self._log_step()
self._record_metrics()
self.step += 1
return SimulationResult(step_records=step_records, summary=self._get_summary())
def _has_work(self, incoming_requests: List[SimRequest]) -> bool:
return bool(incoming_requests) or any(
gpu.pending_requests or gpu.running_requests for gpu in self.gpu_states
)
def _route_requests(self, incoming_requests: List[SimRequest]) -> None:
for req in incoming_requests:
gpu_id = self.router.route(req, self.gpu_states)
self.gpu_states[gpu_id].pending_requests.append(req)
def _schedule_all_gpus(self) -> None:
for gpu in self.gpu_states:
self.scheduler.schedule(gpu)
assert gpu.is_valid(), (
f"GPU{gpu.gpu_id} invalid after scheduling "
f"({gpu.total_seq_len()=}, {gpu.max_total_tokens=})"
)
def _execute_step(self) -> None:
for gpu in self.gpu_states:
gpu.execute_step()
def _log_step(self) -> None:
if self.log_level == 0:
return
parts = [f"step={self.step:<4}"]
for gpu in self.gpu_states:
r, q = len(gpu.running_requests), len(gpu.pending_requests)
if self.log_level == 1:
parts.append(f"GPU{gpu.gpu_id}[R={r:<3} Q={q:<3}]")
else:
run_ids = _format_ids(gpu.running_requests)
queue_ids = _format_ids(gpu.pending_requests)
parts.append(f"GPU{gpu.gpu_id}[R={r}:{run_ids} Q={q}:{queue_ids}]")
print(" | ".join(parts))
def _record_metrics(self) -> None:
for recorder in self.recorders:
recorder.on_step_end(self.step, self.gpu_states)
def _get_summary(self) -> Dict[str, Any]:
return {k: v for r in self.recorders for k, v in r.get_summary().items()}
def _format_ids(requests: List[SimRequest], limit: int = 5) -> str:
if not requests:
return "-"
ids = ",".join(r.request_id for r in requests[:limit])
if len(requests) > limit:
ids += f"...+{len(requests) - limit}"
return ids

View File

@@ -0,0 +1,496 @@
import json
import subprocess
import sys
import tempfile
import unittest
from sglang.srt.debug_utils.schedule_simulator import (
AttentionBalancednessRecorder,
BatchSizeBalancednessRecorder,
FIFOScheduler,
GPUState,
RandomRouter,
RoundRobinRouter,
SimRequest,
SimulationResult,
Simulator,
StepRecord,
generate_random_requests,
load_from_request_logger,
)
from sglang.test.test_utils import CustomTestCase
# ==================== Non-E2E Tests ====================
class TestSimRequest(CustomTestCase):
def test_basic(self):
req = SimRequest(request_id="r1", input_len=100, output_len=50)
self.assertEqual(req.decoded_tokens, 0)
self.assertEqual(req.seq_len(), 100)
self.assertFalse(req.is_finished())
def test_seq_len_with_decoded(self):
req = SimRequest(
request_id="r1", input_len=100, output_len=50, decoded_tokens=10
)
self.assertEqual(req.seq_len(), 110)
def test_is_finished(self):
req = SimRequest(
request_id="r1", input_len=100, output_len=50, decoded_tokens=50
)
self.assertTrue(req.is_finished())
class TestGPUState(CustomTestCase):
def test_batch_size(self):
gpu = GPUState(gpu_id=0, max_total_tokens=10000)
self.assertEqual(gpu.batch_size(), 0)
gpu.running_requests = [
SimRequest(request_id="r1", input_len=100, output_len=50),
SimRequest(request_id="r2", input_len=200, output_len=100),
]
self.assertEqual(gpu.batch_size(), 2)
def test_total_seq_len(self):
gpu = GPUState(gpu_id=0, max_total_tokens=10000)
gpu.running_requests = [
SimRequest(request_id="r1", input_len=100, output_len=50),
SimRequest(
request_id="r2", input_len=200, output_len=100, decoded_tokens=10
),
]
self.assertEqual(gpu.total_seq_len(), 100 + 210)
class TestRouters(CustomTestCase):
def test_round_robin(self):
router = RoundRobinRouter()
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)
results = [router.route(req, gpu_states) for _ in range(8)]
self.assertEqual(results, [0, 1, 2, 3, 0, 1, 2, 3])
def test_random_router(self):
router = RandomRouter()
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)
results = [router.route(req, gpu_states) for _ in range(100)]
self.assertTrue(all(0 <= r < 4 for r in results))
class TestFIFOScheduler(CustomTestCase):
def test_runs_pending_requests(self):
scheduler = FIFOScheduler()
gpu = GPUState(gpu_id=0, max_total_tokens=10000)
gpu.pending_requests = [
SimRequest(request_id=f"r{i}", input_len=100, output_len=50)
for i in range(3)
]
scheduler.schedule(gpu)
self.assertEqual(len(gpu.running_requests), 3)
self.assertEqual(len(gpu.pending_requests), 0)
def test_respects_token_limit(self):
scheduler = FIFOScheduler()
gpu = GPUState(gpu_id=0, max_total_tokens=250)
gpu.pending_requests = [
SimRequest(request_id=f"r{i}", input_len=100, output_len=50)
for i in range(5)
]
scheduler.schedule(gpu)
self.assertEqual(len(gpu.running_requests), 2)
self.assertEqual(len(gpu.pending_requests), 3)
def test_evicts_lifo_when_over_budget(self):
scheduler = FIFOScheduler()
gpu = GPUState(gpu_id=0, max_total_tokens=250)
gpu.running_requests = [
SimRequest(request_id=f"r{i}", input_len=100, output_len=50)
for i in range(3)
] # 300 tokens total
scheduler.schedule(gpu)
self.assertEqual(len(gpu.running_requests), 2)
self.assertEqual(len(gpu.pending_requests), 1)
self.assertEqual(gpu.pending_requests[0].request_id, "r2")
class TestMetrics(CustomTestCase):
def test_batch_size_balancedness(self):
recorder = BatchSizeBalancednessRecorder()
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(2)]
gpu_states[0].running_requests = [
SimRequest(request_id="r1", input_len=100, output_len=50)
]
gpu_states[1].running_requests = [
SimRequest(request_id="r2", input_len=100, output_len=50),
SimRequest(request_id="r3", input_len=100, output_len=50),
]
recorder.on_step_end(0, gpu_states)
self.assertAlmostEqual(
recorder.get_summary()["batch_size_balancedness_mean"], 0.75
)
def test_attention_balancedness(self):
recorder = AttentionBalancednessRecorder()
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(2)]
gpu_states[0].running_requests = [
SimRequest(request_id="r1", input_len=100, output_len=50)
]
gpu_states[1].running_requests = [
SimRequest(request_id="r2", input_len=200, output_len=50)
]
recorder.on_step_end(0, gpu_states)
self.assertAlmostEqual(
recorder.get_summary()["attention_balancedness_mean"], 0.75
)
def test_empty_history(self):
recorder = BatchSizeBalancednessRecorder()
self.assertEqual(recorder.get_summary()["batch_size_balancedness_mean"], 0.0)
def test_all_zero_batch_size(self):
recorder = BatchSizeBalancednessRecorder()
gpu_states = [GPUState(gpu_id=i, max_total_tokens=10000) for i in range(2)]
recorder.on_step_end(0, gpu_states)
self.assertAlmostEqual(
recorder.get_summary()["batch_size_balancedness_mean"], 1.0
)
class TestDataLoader(CustomTestCase):
def test_load_from_request_logger(self):
log_data = [
{"event": "request.received", "rid": "r1", "obj": {"text": "hello"}},
{
"event": "request.finished",
"rid": "r1",
"out": {"meta_info": {"prompt_tokens": 100, "completion_tokens": 50}},
},
{
"event": "request.finished",
"rid": "r2",
"out": {"meta_info": {"prompt_tokens": 200, "completion_tokens": 100}},
},
]
with tempfile.NamedTemporaryFile(mode="w", suffix=".log", delete=False) as f:
for item in log_data:
f.write(json.dumps(item) + "\n")
f.flush()
requests = load_from_request_logger(f.name)
self.assertEqual(len(requests), 2)
self.assertEqual(requests[0].request_id, "r1")
self.assertEqual(requests[0].input_len, 100)
self.assertEqual(requests[1].input_len, 200)
def test_empty_file(self):
with tempfile.NamedTemporaryFile(mode="w", suffix=".log", delete=False) as f:
f.write("")
f.flush()
self.assertEqual(len(load_from_request_logger(f.name)), 0)
class TestDataSynthesis(CustomTestCase):
def test_generate_basic(self):
requests = generate_random_requests(
num_requests=10, input_len=100, output_len=50
)
self.assertEqual(len(requests), 10)
for req in requests:
self.assertEqual(req.input_len, 100)
self.assertEqual(req.output_len, 50)
def test_generate_with_range_ratio(self):
requests = generate_random_requests(
num_requests=100, input_len=100, output_len=50, range_ratio=0.5, seed=42
)
for req in requests:
self.assertGreaterEqual(req.input_len, 50)
self.assertLessEqual(req.input_len, 100)
def test_generate_with_seed(self):
r1 = generate_random_requests(
num_requests=10, input_len=100, output_len=50, range_ratio=0.5, seed=42
)
r2 = generate_random_requests(
num_requests=10, input_len=100, output_len=50, range_ratio=0.5, seed=42
)
for a, b in zip(r1, r2):
self.assertEqual(a.input_len, b.input_len)
class TestSimulator(CustomTestCase):
def test_basic_run(self):
requests = [
SimRequest(request_id=f"r{i}", input_len=10, output_len=5)
for i in range(10)
]
sim = Simulator(
num_gpus=2,
router=RoundRobinRouter(),
scheduler=FIFOScheduler(),
recorders=[
BatchSizeBalancednessRecorder(),
AttentionBalancednessRecorder(),
],
max_total_tokens=100,
)
result = sim.run(requests)
self.assertIsInstance(result, SimulationResult)
self.assertIn("batch_size_balancedness_mean", result.summary)
self.assertGreater(len(result.step_records), 0)
def test_all_requests_complete(self):
requests = [
SimRequest(request_id=f"r{i}", input_len=10, output_len=3) for i in range(4)
]
sim = Simulator(
num_gpus=2,
router=RoundRobinRouter(),
scheduler=FIFOScheduler(),
max_total_tokens=10000,
)
sim.run(requests)
for gpu in sim.gpu_states:
self.assertEqual(len(gpu.pending_requests), 0)
self.assertEqual(len(gpu.running_requests), 0)
def test_empty_requests(self):
sim = Simulator(
num_gpus=2, router=RoundRobinRouter(), scheduler=FIFOScheduler()
)
result = sim.run([])
self.assertEqual(result.summary, {})
self.assertEqual(len(result.step_records), 0)
def test_step_records(self):
requests = [
SimRequest(request_id=f"r{i}", input_len=10, output_len=3) for i in range(4)
]
sim = Simulator(
num_gpus=2,
router=RoundRobinRouter(),
scheduler=FIFOScheduler(),
max_total_tokens=10000,
)
result = sim.run(requests)
self.assertGreater(len(result.step_records), 0)
for record in result.step_records:
self.assertIsInstance(record, StepRecord)
self.assertIn(record.gpu_id, [0, 1])
self.assertEqual(len([r for r in result.step_records if r.step == 0]), 2)
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 = [
SimRequest(request_id="r0", input_len=50, output_len=10),
SimRequest(request_id="r1", input_len=50, output_len=10),
]
sim = Simulator(
num_gpus=1,
router=RoundRobinRouter(),
scheduler=FIFOScheduler(),
max_total_tokens=110,
)
result = sim.run(requests)
# Check that preemption happened at some point
found_preemption = False
for record in result.step_records:
if record.running_count == 1 and record.pending_count == 1:
found_preemption = True
break
self.assertTrue(
found_preemption, "Expected preemption to occur due to token growth"
)
# ==================== E2E Tests ====================
class TestCLI(CustomTestCase):
def _run_cli(self, *args):
return subprocess.run(
[sys.executable, "-m", "sglang.srt.debug_utils.schedule_simulator", *args],
capture_output=True,
text=True,
)
def _assert_output_contains(self, output: str, expected_lines: str):
for line in expected_lines.strip().split("\n"):
self.assertIn(line, output)
def test_cli_basic(self):
log_data = [
{
"event": "request.finished",
"rid": "r1",
"out": {"meta_info": {"prompt_tokens": 100, "completion_tokens": 50}},
},
{
"event": "request.finished",
"rid": "r2",
"out": {"meta_info": {"prompt_tokens": 200, "completion_tokens": 100}},
},
]
with tempfile.NamedTemporaryFile(mode="w", suffix=".log", delete=False) as f:
for item in log_data:
f.write(json.dumps(item) + "\n")
input_file = f.name
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
output_file = f.name
result = self._run_cli(
"--input", input_file, "--num-gpus", "2", "--output", output_file
)
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self.assertIn("Loaded 2 requests", result.stdout)
with open(output_file) as f:
self.assertIn("batch_size_balancedness_mean", json.load(f))
def test_cli_random_router(self):
log_data = [
{
"event": "request.finished",
"rid": "r1",
"out": {"meta_info": {"prompt_tokens": 100, "completion_tokens": 50}},
}
]
with tempfile.NamedTemporaryFile(mode="w", suffix=".log", delete=False) as f:
for item in log_data:
f.write(json.dumps(item) + "\n")
input_file = f.name
result = self._run_cli("--input", input_file, "--router", "random")
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self.assertIn("router=random", result.stdout)
def test_cli_synthetic(self):
result = self._run_cli(
"--synthetic",
"--synth-num-requests",
"100",
"--synth-input-len",
"512",
"--synth-output-len",
"128",
"--synth-range-ratio",
"0.5",
"--num-gpus",
"4",
)
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self.assertIn("Generated 100 synthetic requests", result.stdout)
def test_cli_log_level(self):
result = self._run_cli(
"--synthetic",
"--synth-num-requests",
"10",
"--synth-output-len",
"5",
"--num-gpus",
"2",
"--log-level",
"1",
)
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self.assertIn("step=", result.stdout)
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(
"--synthetic",
"--synth-num-requests",
"4",
"--synth-input-len",
"10",
"--synth-output-len",
"2",
"--synth-seed",
"42",
"--num-gpus",
"2",
"--max-total-tokens",
"10000",
"--log-level",
"2",
)
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self.assertIn(
"step=0 | GPU0[R=2:syn0,syn2 Q=0:-] | GPU1[R=2:syn1,syn3 Q=0:-]",
result.stdout,
)
self.assertIn(
"step=1 | GPU0[R=0:- Q=0:-] | GPU1[R=0:- Q=0:-]", result.stdout
)
self.assertIn("batch_size_balancedness_mean: 1.0000", result.stdout)
def test_e2e_queuing_due_to_token_limit(self):
result = self._run_cli(
"--synthetic",
"--synth-num-requests",
"4",
"--synth-input-len",
"100",
"--synth-output-len",
"3",
"--synth-seed",
"42",
"--num-gpus",
"1",
"--max-total-tokens",
"210",
"--log-level",
"2",
)
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self._assert_output_contains(
result.stdout,
"""
step=0 | GPU0[R=2:syn0,syn1 Q=2:syn2,syn3]
step=1 | GPU0[R=2:syn0,syn1 Q=2:syn2,syn3]
step=2 | GPU0[R=0:- Q=2:syn2,syn3]
step=3 | GPU0[R=2:syn2,syn3 Q=0:-]
step=4 | GPU0[R=2:syn2,syn3 Q=0:-]
step=5 | GPU0[R=0:- Q=0:-]""",
)
def test_e2e_retraction_due_to_token_growth(self):
result = self._run_cli(
"--synthetic",
"--synth-num-requests",
"2",
"--synth-input-len",
"50",
"--synth-output-len",
"10",
"--synth-seed",
"42",
"--num-gpus",
"1",
"--max-total-tokens",
"110",
"--log-level",
"2",
)
self.assertEqual(result.returncode, 0, f"CLI failed: {result.stderr}")
self._assert_output_contains(
result.stdout,
"""
step=0 | GPU0[R=2:syn0,syn1 Q=0:-]
step=5 | GPU0[R=2:syn0,syn1 Q=0:-]
step=6 | GPU0[R=1:syn0 Q=1:syn1]
step=9 | GPU0[R=0:- Q=1:syn1]
step=10 | GPU0[R=1:syn1 Q=0:-]
step=13 | GPU0[R=0:- Q=0:-]""",
)
if __name__ == "__main__":
unittest.main()