[Simulator] Add high-fidelity CPU-based inference simulator (#33824)

Co-authored-by: zhouhaizhu.zhz <zhouhaizhu.zhz@alibaba-inc.com>
Co-authored-by: LinSiyuan814 <linsiyuan.lsy@alibaba-inc.com>
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
Ruiyan Ma
2026-09-04 11:12:11 +08:00
committed by GitHub
co-authored by zhouhaizhu.zhz LinSiyuan814 hzh0425
parent a5f07b1241
commit 59799a3687
81 changed files with 6624 additions and 0 deletions
@@ -19,6 +19,8 @@ on:
outputs:
main_package:
value: ${{ jobs.run.outputs.main_package }}
simulator:
value: ${{ jobs.run.outputs.simulator }}
sgl_kernel:
value: ${{ jobs.run.outputs.sgl_kernel }}
jit_kernel:
@@ -44,6 +46,7 @@ jobs:
runs-on: ubuntu-latest
outputs:
main_package: ${{ steps.filter.outputs.main_package || steps.run-mode.outputs.run_all_tests }}
simulator: ${{ steps.filter.outputs.simulator || steps.run-mode.outputs.run_all_tests }}
sgl_kernel: ${{ steps.filter.outputs.sgl_kernel }}
jit_kernel: ${{ steps.filter.outputs.jit_kernel || steps.run-mode.outputs.run_all_tests }}
multimodal_gen: ${{ steps.filter.outputs.multimodal_gen || steps.run-mode.outputs.run_all_tests }}
@@ -94,6 +97,11 @@ jobs:
- "test/**/!(*.md)"
- "rust/**"
- "proto/sglang/runtime/v1/sglang.proto"
simulator:
- ".github/workflows/pr-test-extra.yml"
- ".github/workflows/_pr-test-simulator-cpu.yml"
- "benchmark/simulator/**/!(*.md)"
- "tools/sglang-simulator/**/!(*.md)"
multimodal_gen:
- ".github/workflows/pr-test.yml"
- ".github/workflows/pr-test-multimodal-gen.yml"
@@ -0,0 +1,96 @@
name: PR Test SGLang Simulator (CPU)
on:
workflow_call:
inputs:
check_changes:
description: 'toJson(needs.check-changes.outputs).'
type: string
required: true
caller_inputs:
description: 'toJson(inputs) from the caller workflow.'
type: string
required: true
rust_ext_artifact:
description: 'Artifact of prebuilt Rust extension modules.'
type: string
default: ''
env:
SGLANG_IS_IN_CI: true
SKIP_PR_TEST_HEALTH_CHECK: ${{ (fromJson(inputs.caller_inputs).skip_pr_test_health_check || fromJson(inputs.caller_inputs).test_parallel_dispatch || fromJson(inputs.caller_inputs).run_all_tests) && 'true' || 'false' }}
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
USE_VENV: false
jobs:
run:
name: simulator-test-cpu
if: fromJson(inputs.check_changes).main_package == 'true' || fromJson(inputs.check_changes).simulator == 'true'
runs-on: ubuntu-latest
timeout-minutes: 40
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
ref: ${{ fromJson(inputs.caller_inputs).git_ref || github.sha }}
- uses: ./.github/actions/check-pr-test-health
- uses: ./.github/actions/check-maintenance
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Install uv
uses: astral-sh/setup-uv@v5
- uses: ./.github/actions/download-rust-ext
id: rust_ext
with:
artifact_name: ${{ inputs.rust_ext_artifact }}
- name: Install protoc + Rust toolchain
if: ${{ steps.rust_ext.outputs.hit != 'true' }}
timeout-minutes: 10
run: bash scripts/ci/utils/install_rust_protoc.sh
- name: Rust cache (rust/ workspace)
if: ${{ steps.rust_ext.outputs.hit != 'true' }}
uses: Swatinem/rust-cache@v2
with:
workspaces: rust
shared-key: "sglang-grpc-cpu"
- name: Install dependencies
timeout-minutes: 20
env:
UV_SYSTEM_PYTHON: "1"
run: |
uv pip install -e "python" --index-strategy unsafe-best-match --prerelease allow
uv pip install pytest
uv pip install -e "tools/sglang-simulator[aic]"
- name: Prebuild HiCache native hash extension
timeout-minutes: 5
env:
MALLOC_ARENA_MAX: "2"
MAX_JOBS: "1"
TORCH_EXTENSIONS_DIR: ${{ runner.temp }}/torch-extensions
run: |
python3 -c \
'from sglang.srt.mem_cache.cpp_utils.native_hash import get_native_hash; get_native_hash([1], None)'
- name: Run SGLang Simulator compatibility tests
timeout-minutes: 10
env:
MALLOC_ARENA_MAX: "2"
MAX_JOBS: "1"
PYTEST_DISABLE_PLUGIN_AUTOLOAD: "1"
TORCH_EXTENSIONS_DIR: ${{ runner.temp }}/torch-extensions
run: |
python3 -m pytest -q tools/sglang-simulator/test/test_simulation_sglang_runner.py
python3 -m pytest -q tools/sglang-simulator/test/test_simulation_sglang_serving.py
python3 -m pytest -q tools/sglang-simulator/test/test_simulation_cache_hit_ratio.py
python3 -m pytest -q tools/sglang-simulator/test/test_simulation_offline_blocking.py
+11
View File
@@ -149,6 +149,16 @@ jobs:
skip_pr_test_health_check: ${{ inputs.skip_pr_test_health_check == true }}
secrets: inherit
simulator-test-cpu:
needs: [check-changes, call-gate, rust-ext-build]
if: ${{ !failure() && !cancelled() && needs.check-changes.result == 'success' && (needs.call-gate.result == 'success' || needs.call-gate.result == 'skipped') }}
uses: ./.github/workflows/_pr-test-simulator-cpu.yml
with:
check_changes: ${{ toJson(needs.check-changes.outputs) }}
caller_inputs: ${{ toJson(inputs) }}
rust_ext_artifact: ${{ needs.rust-ext-build.outputs.artifact_name }}
# No `secrets: inherit`: this hosted CPU job has no secret consumer.
# =============================================== extra-a (1-/2-gpu) ===============================================
extra-a-test-1-gpu-small:
needs: [check-changes, call-gate, sgl-kernel-build-wheels, rust-ext-build]
@@ -249,6 +259,7 @@ jobs:
call-gate,
sgl-kernel-build-wheels,
rust-ext-build,
simulator-test-cpu,
extra-a-test-1-gpu-small,
extra-a-test-1-gpu-large,
extra-a-test-2-gpu-large,
+1
View File
@@ -171,6 +171,7 @@ benchmark/mmlu/data.tar
benchmark/llava_bench/images
benchmark/llava_bench/mme_pack
*.jsonl
!tools/sglang-simulator/examples/replay/trace.jsonl
tmp*.txt
/tmp/
+168
View File
@@ -0,0 +1,168 @@
"""In-process benchmark runner for SGLang Simulator."""
import asyncio
import atexit
import json
import os
from dataclasses import asdict
from typing import Iterator
import numpy as np
from sglang_simulator.compat import apply_simulator_server_args
from sglang_simulator.dataset import BaseDataset, GenericRequest
from sglang_simulator.simulation.benchmark import BaseBenchmarkRunner, BenchmarkConfig
from sglang_simulator.utils.logger import get_logger
SGLANG_SIMULATOR_OUTPUT_DIR = os.getenv(
"SGLANG_SIMULATOR_OUTPUT_DIR", "/tmp/sglang_simulator/output"
)
SIMULATION_METRICS_PATH = f"{SGLANG_SIMULATOR_OUTPUT_DIR}/metrics.json"
os.environ["SGLANG_SIMULATOR_OUTPUT_DIR"] = SGLANG_SIMULATOR_OUTPUT_DIR
if os.getenv("SGLANG_SIMULATOR_OUTPUT_MODE") is None:
os.environ["SGLANG_SIMULATOR_OUTPUT_MODE"] = "OFFLINE"
# Import the simulator engine only after configuring its worker environment.
from sglang_simulator.simulation.sglang.engine import ( # noqa: E402
SGLangSimulationEngine,
)
# SGLang must be imported after the simulator hooks are installed by engine.py.
from sglang.srt.server_args import ServerArgs # noqa: E402
logger = get_logger("sglang_simulator")
class SGLangBenchmarkRunner(BaseBenchmarkRunner):
"""Run a simulator workload directly through SGLang's in-process Engine."""
def __init__(self, server_args: ServerArgs):
# Disable features that are unnecessary for simulation.
server_args_kwargs = asdict(server_args)
apply_simulator_server_args(server_args_kwargs)
self.engine = SGLangSimulationEngine(**server_args_kwargs)
self.server_args = self.engine.server_args
self._shutdown = False
def flush_cache(self):
self.engine.flush_cache()
def clear_hicache_storage(self):
self.engine.loop.run_until_complete(
self.engine.tokenizer_manager.clear_hicache_storage()
)
def get_request(
self,
dataset: BaseDataset,
ignore_timestamp: bool = False,
request_rate: float = float("inf"),
) -> Iterator[tuple[GenericRequest, dict]]:
yield_delay = 0
for req in dataset:
if ignore_timestamp:
created_time = yield_delay
yield_delay += np.random.exponential(1.0 / request_rate)
else:
created_time = req.custom_params.get("created_time", 0)
simulation_params = {
"total_request": len(dataset), # Include the warmup requests.
"created_time": created_time,
}
yield (req, simulation_params)
async def async_benchmark(
self,
benchmark_config: BenchmarkConfig,
dataset: BaseDataset,
):
await self.engine.tokenizer_manager.start_profile()
if os.path.exists(SIMULATION_METRICS_PATH):
with open(SIMULATION_METRICS_PATH, "w") as metrics_file:
# Clear data from a previous benchmark in the same process.
pass
tasks = []
logger.info(f"Created {len(dataset)} request tasks.")
for req, simulation_params in self.get_request(
dataset,
ignore_timestamp=benchmark_config.ignore_request_timestamp,
request_rate=benchmark_config.request_rate,
):
task = asyncio.create_task(
self.engine.async_generate(
prompt=req.prompt,
input_ids=req.token_ids,
sampling_params={
"ignore_eos": True,
"max_new_tokens": req.output_length,
"custom_params": {
# Transfer simulation arguments through sampling params.
"simulation": simulation_params
},
},
)
)
tasks.append(task)
_ = await asyncio.gather(*tasks)
# Trigger the simulator's profile handler to flush final metrics.
await self.engine.tokenizer_manager.start_profile()
if os.path.exists(SIMULATION_METRICS_PATH):
with open(SIMULATION_METRICS_PATH) as metrics_file:
metrics = json.load(metrics_file)
else:
logger.error(
f"Failed to load metrics from serving backend. The metrics file "
f"should be loaded from {SIMULATION_METRICS_PATH}."
)
return None
return metrics
def benchmark(self, benchmark_config: BenchmarkConfig, dataset: BaseDataset):
return self.engine.loop.run_until_complete(
self.async_benchmark(benchmark_config, dataset)
)
def get_iteration_stats(self) -> list[dict]:
data = []
file_path = f"{SGLANG_SIMULATOR_OUTPUT_DIR}/iteration.jsonl"
if os.path.exists(file_path):
with open(file_path) as stats_file:
line = stats_file.readline()
while line:
data.append(json.loads(line))
line = stats_file.readline()
else:
logger.error(f"The iteration statistics data({file_path}) does not exist.")
return data
def get_request_stats(self) -> list[dict]:
data = []
file_path = f"{SGLANG_SIMULATOR_OUTPUT_DIR}/request.jsonl"
if os.path.exists(file_path):
with open(file_path) as stats_file:
line = stats_file.readline()
while line:
data.append(json.loads(line))
line = stats_file.readline()
else:
logger.error(f"The request statistics data({file_path}) does not exist.")
return data
def shutdown(self):
if self._shutdown:
return None
logger.info("Attempting to shut down the SGLang backend engine.")
try:
return self.engine.shutdown()
finally:
self._shutdown = True
atexit.unregister(self.engine.shutdown)
+280
View File
@@ -0,0 +1,280 @@
"""SGLang serving benchmark adapter for simulator traffic.
This script deliberately reuses SGLang's benchmark implementation and dataset
loaders. It only owns the simulator-specific parts of the protocol:
* convert request-rate or trace timestamps into logical arrival timestamps;
* inject the internal ``sampling_params.custom_params.simulation`` metadata;
* avoid client-side pacing in OFFLINE mode; and
* display backend-produced simulator metrics when they are locally available.
User datasets must not contain simulator metadata.
"""
import argparse
import contextlib
import json
import os
import re
import sys
from dataclasses import fields
from pathlib import Path
from typing import AsyncGenerator, List, Optional
import aiohttp
import numpy as np
from sglang_simulator.compat import validate_benchmark_runtime
from sglang_simulator.dataset.autobench import register_autobench_dataset
register_autobench_dataset()
from sglang.benchmark import serving
from sglang.benchmark.datasets.common import DatasetRow
_ORIGINAL_AIOHTTP_REQUEST = None
_ORIGINAL_CALCULATE_METRICS = serving.calculate_metrics
_ORIGINAL_GET_REQUEST = serving.get_request
_ORIGINAL_RUN_BENCHMARK = serving.run_benchmark
_SIMULATOR_MODE = "offline"
_USE_TRACE_TIMESTAMPS = False
def _metrics_path() -> Path:
output_dir = Path(
os.getenv("SGLANG_SIMULATOR_OUTPUT_DIR", "/tmp/sglang_simulator/output")
)
return output_dir / "metrics.json"
def _load_backend_metrics() -> Optional[dict]:
metrics_path = _metrics_path()
if not metrics_path.is_file():
return None
return json.loads(metrics_path.read_text(encoding="utf-8"))
class _DurationReplacingStream:
"""Keep SGLang's output format but print the simulated duration."""
def __init__(self, target):
self.target = target
def write(self, text):
if "Benchmark duration (s):" in text:
metrics = _load_backend_metrics()
if metrics is not None and "duration" in metrics:
text = "{:<40} {:<10.2f}".format(
"Benchmark duration (s):", metrics["duration"]
)
return self.target.write(text)
def flush(self):
return self.target.flush()
def _set_simulation_metadata(
request: DatasetRow, *, created_time_ms: float, total_request: int
) -> None:
"""Attach transient metadata without replacing dataset-specific parameters."""
extra_request_body = dict(request.extra_request_body or {})
extra_request_body["simulation"] = {
"created_time_ms": created_time_ms,
"total_request": total_request,
}
request.extra_request_body = extra_request_body
async def simulator_get_request(
input_requests: List[DatasetRow],
request_rate: float,
use_trace_timestamps: bool = False,
slowdown_factor: float = 1.0,
) -> AsyncGenerator[DatasetRow, None]:
"""Generate simulator traffic while retaining official BLOCKING pacing."""
# The benchmark may not forward --use-trace-timestamps to get_request(),
# so preserve the parsed value in this adapter.
use_trace_timestamps = use_trace_timestamps or _USE_TRACE_TIMESTAMPS
if _SIMULATOR_MODE == "blocking":
async for request in _ORIGINAL_GET_REQUEST(
input_requests,
request_rate,
use_trace_timestamps=use_trace_timestamps,
slowdown_factor=slowdown_factor,
):
yield request
return
total_request = len(input_requests)
if use_trace_timestamps:
if any(request.timestamp is None for request in input_requests):
raise ValueError(
"--use-trace-timestamps requires every request to have timestamp"
)
input_requests.sort(key=lambda request: request.timestamp)
trace_start_time_ms = input_requests[0].timestamp if input_requests else 0.0
for request in input_requests:
created_time_ms = (
float(request.timestamp) - float(trace_start_time_ms)
) * slowdown_factor
_set_simulation_metadata(
request,
created_time_ms=created_time_ms,
total_request=total_request,
)
yield request
return
created_time_ms = 0.0
for request in input_requests:
_set_simulation_metadata(
request,
created_time_ms=created_time_ms,
total_request=total_request,
)
yield request
if request_rate != float("inf"):
created_time_ms += np.random.exponential(1.0 / request_rate) * 1000.0
def install_aiohttp_json_hijack(
*, hijack_url_regex: Optional[str] = r"/generate(?:\?.*)?$"
) -> None:
"""Move transient metadata into the already-built sampling parameters."""
global _ORIGINAL_AIOHTTP_REQUEST
if _ORIGINAL_AIOHTTP_REQUEST is not None:
return
pattern = re.compile(hijack_url_regex) if hijack_url_regex else None
_ORIGINAL_AIOHTTP_REQUEST = aiohttp.ClientSession._request
async def patched_request(self, method, url, **kwargs):
if pattern is None or pattern.search(str(url)):
payload = kwargs.get("json")
if isinstance(payload, dict) and "simulation" in payload:
simulation = payload.pop("simulation")
sampling_params = payload.setdefault("sampling_params", {})
custom_params = sampling_params.setdefault("custom_params", {})
custom_params["simulation"] = simulation
kwargs["json"] = payload
return await _ORIGINAL_AIOHTTP_REQUEST(self, method, url, **kwargs)
aiohttp.ClientSession._request = patched_request
def simulator_calculate_metrics(*args, **kwargs):
"""Use simulator metrics; mark unsupported client-only fields with -1."""
client_metrics, output_lens = _ORIGINAL_CALCULATE_METRICS(*args, **kwargs)
backend_metrics = _load_backend_metrics()
if backend_metrics is None:
print(
f"Simulator metrics are not available at {_metrics_path()}; "
"showing client-side benchmark metrics."
)
return client_metrics, output_lens
metric_names = {field.name for field in fields(serving.BenchmarkMetrics)}
values = {name: backend_metrics.get(name, -1) for name in metric_names}
return serving.BenchmarkMetrics(**values), output_lens
def _replace_output_file_duration(
args: argparse.Namespace, simulated_duration: float
) -> None:
output_file = getattr(args, "output_file", None)
if not output_file:
return
path = Path(output_file)
if not path.is_file():
return
lines = path.read_text(encoding="utf-8").splitlines()
if not lines:
return
last_result = json.loads(lines[-1])
last_result["duration"] = simulated_duration
lines[-1] = json.dumps(last_result)
path.write_text("\n".join(lines) + "\n", encoding="utf-8")
def simulator_run_benchmark(args: argparse.Namespace):
global _USE_TRACE_TIMESTAMPS
if args.backend != "sglang":
raise ValueError(
"benchmark/simulator/bench_serving.py requires --backend sglang"
)
if args.dataset_name == "mooncake":
raise ValueError(
"Mooncake's multi-round scheduler is not supported by the simulator "
"benchmark adapter"
)
_USE_TRACE_TIMESTAMPS = getattr(args, "use_trace_timestamps", False)
args.profile = True
with contextlib.redirect_stdout(_DurationReplacingStream(sys.stdout)):
result = _ORIGINAL_RUN_BENCHMARK(args)
backend_metrics = _load_backend_metrics()
if backend_metrics is not None and "duration" in backend_metrics:
simulated_duration = backend_metrics["duration"]
if isinstance(result, dict):
result["duration"] = simulated_duration
_replace_output_file_duration(args, simulated_duration)
return result
def _extract_simulator_args(argv: list[str]) -> tuple[str, list[str]]:
parser = argparse.ArgumentParser(add_help=False)
parser.add_argument(
"--simulator-mode",
choices=("offline", "blocking"),
default="offline",
help=argparse.SUPPRESS,
)
args, remaining = parser.parse_known_args(argv)
return args.simulator_mode, remaining
def _simulator_argument_parser(base_parser):
"""Include simulator-owned datasets in SGLang's hard-coded CLI choices."""
class SimulatorArgumentParser(base_parser):
def add_argument(self, *name_or_flags, **kwargs):
choices = kwargs.get("choices")
if (
"--dataset-name" in name_or_flags
and choices is not None
and "autobench" not in choices
):
kwargs["choices"] = [*choices, "autobench"]
if "--warmup-requests" in name_or_flags:
kwargs["default"] = 0
return super().add_argument(*name_or_flags, **kwargs)
return SimulatorArgumentParser
def cli_main() -> None:
global _SIMULATOR_MODE
validate_benchmark_runtime()
if any(argument in ("-h", "--help") for argument in sys.argv[1:]):
print(
"SGLang Simulator option: "
"--simulator-mode {offline,blocking} (default: offline)\n"
)
_SIMULATOR_MODE, remaining = _extract_simulator_args(sys.argv[1:])
sys.argv = [sys.argv[0], *remaining]
serving.get_request = simulator_get_request
serving.calculate_metrics = simulator_calculate_metrics
serving.run_benchmark = simulator_run_benchmark
install_aiohttp_json_hijack()
print(f"SGLang Simulator benchmark mode: {_SIMULATOR_MODE.upper()}")
original_parser = serving.ArgumentParser
serving.ArgumentParser = _simulator_argument_parser(original_parser)
try:
serving.cli_main()
finally:
serving.ArgumentParser = original_parser
if __name__ == "__main__":
cli_main()
+1
View File
@@ -953,6 +953,7 @@
"docs/advanced_features/hicache_storage_runtime_attach_detach"
]
},
"docs/advanced_features/sglang_simulator",
"docs/advanced_features/vlm_query",
"docs/advanced_features/dp_for_multi_modal_encoder",
"docs/advanced_features/cuda_graph_for_multi_modal_encoder",
+1
View File
@@ -16,5 +16,6 @@ description: Advanced configuration, optimization, and deployment features for S
- [PD Disaggregation](./pd_disaggregation)
- [Pipeline Parallelism](./pipeline_parallelism)
- [HiCache](./hicache_best_practices)
- [SGLang Simulator](./sglang_simulator)
- [Observability](./observability)
- [And more…](./server_arguments)
@@ -0,0 +1,88 @@
---
title: "SGLang Simulator"
metatags:
description: "Run SGLang scheduling and KV-cache simulations without loading model weights or executing GPU kernels."
---
SGLang Simulator reuses SGLang's scheduler, request lifecycle, and KV-cache implementation while replacing model forward execution with a latency predictor. Use it to compare scheduling and cache configurations on timestamped or synthetic workloads without loading model weights.
## Supported scope
SGLang Simulator tracks the current `main` branch and recent SGLang releases. The current integration is validated with `v0.5.16`, `v0.5.17`, `v0.5.18`, and `main`.
The initial upstream scope uses one simulated worker with `tp_size=1`, `ep_size=1`, `dp_size=1`, and `pp_size=1`. A simulator configuration can describe a larger target system for latency prediction, but the SGLang runtime process topology remains single-worker.
The simulator supports:
- synthetic request rates, ShareGPT workloads, and timestamped Autobench traces;
- OFFLINE logical-time simulation and BLOCKING wall-clock replay;
- AIConfigurator, ML, and replay latency predictors;
- SGLang prefix caching and [HiCache](/docs/advanced_features/hicache); and
- serving-compatible TTFT, TPOT, ITL, throughput, and cache-hit metrics.
## Install from the SGLang repository
Use the simulator and SGLang source from the same monorepo checkout:
```bash
python3 -m pip install -e tools/sglang-simulator
export PYTHONPATH="$PWD/tools/sglang-simulator/src:$PWD/python"
```
AIConfigurator is optional. Install the validated extra only when you use an AIConfigurator predictor:
```bash
python3 -m pip install -e "tools/sglang-simulator[aic]"
```
## Start a simulator server
Choose a fresh output directory for every run. The server owns the simulation mode and writes metrics to this directory.
```bash
export SGLANG_USE_CPU_ENGINE=1
export CUDA_VISIBLE_DEVICES=""
export SGLANG_SIMULATOR_OUTPUT_MODE=OFFLINE
export SGLANG_SIMULATOR_OUTPUT_DIR=/tmp/sglang-simulator-quickstart
python3 -m sglang_simulator.simulation.sglang.launch_server \
--model-path tools/sglang-simulator/test/assets/qwen3-8b \
--tokenizer-path tools/sglang-simulator/examples/assets/tokenizer \
--sim-config-path tools/sglang-simulator/examples/sim_configs/replay.json \
--port 30000
```
`OFFLINE` advances the simulator's logical clock without sleeping. `BLOCKING` also sleeps for predicted forward and cache-load latency, which is useful when a client must observe simulated wall-clock pacing.
## Send a workload
In another terminal, export the same output directory and run the simulator-aware serving benchmark from the repository root:
```bash
export PYTHONPATH="$PWD/tools/sglang-simulator/src:$PWD/python"
export SGLANG_SIMULATOR_OUTPUT_DIR=/tmp/sglang-simulator-quickstart
python3 benchmark/simulator/bench_serving.py \
--simulator-mode offline \
--backend sglang \
--base-url http://127.0.0.1:30000 \
--model tools/sglang-simulator/test/assets/qwen3-8b \
--tokenizer tools/sglang-simulator/examples/assets/tokenizer \
--dataset-name sharegpt \
--dataset-path tools/sglang-simulator/examples/workloads/sharegpt-example.json \
--sharegpt-output-len 4 \
--num-prompts 3 \
--output-file /tmp/sglang-simulator-quickstart/benchmark.json
```
The benchmark injects logical arrival metadata into each request and displays the server-side simulator metrics. For timestamped traffic, use the simulator-owned Autobench JSONL format and add `--use-trace-timestamps`.
## Read the results
The output directory contains:
- `metrics.json`: aggregate latency, throughput, and cache metrics;
- `request.jsonl`: per-request timing and cache information; and
- `iteration.jsonl`: scheduler batch composition and predicted iteration latency.
Use a unique output directory for each run so metrics from separate experiments are not mixed. See the [SGLang Simulator source README](https://github.com/sgl-project/sglang/tree/main/tools/sglang-simulator) for simulator configuration fields, predictor examples, and maintained tests.
+258
View File
@@ -0,0 +1,258 @@
# SGLang Simulator
SGLang Simulator reuses SGLang's scheduler and cache implementation while
replacing model forward execution with a latency predictor. It supports
timestamped trace replay, synthetic workloads, hierarchical cache simulation,
and serving-compatible metrics without loading model weights.
See the [SGLang Simulator advanced-feature guide](../../docs/docs/advanced_features/sglang_simulator.mdx)
for the user-facing setup and serving workflow.
## Compatibility
SGLang Simulator tracks the current SGLang `main` branch and maintains compatibility
with recent SGLang releases. The current integration is validated with `v0.5.16`,
`v0.5.17`, `v0.5.18`, and `main`. Compatibility code uses API and capability
checks instead of branching on version numbers.
## Requirements
- A compatible SGLang checkout. The simulator uses the SGLang source from the
same monorepo checkout.
- A local model directory containing model configuration files. Tokenizer files
are also required unless tokenizer initialization is disabled.
- Predictor data for AIConfigurator, ML, or replay mode.
Use an official SGLang image matching the checkout when validating GPU and
runtime compatibility.
## Installation
From the SGLang repository:
```bash
pip install -e tools/sglang-simulator
```
The simulator does not install or pin a second `sglang` package. Run it from a
checkout whose `python/sglang` package is available on `PYTHONPATH`, or from a
matching official SGLang image.
AIConfigurator is optional. Install it separately when using the
`aiconfigurator` predictor. The `aic` extra pins AIConfigurator to the exact release
validated with the simulator so upstream API changes cannot silently alter an
installation. In a clean virtual environment, install the extra with:
```bash
pip install -e "tools/sglang-simulator[aic]"
```
In an existing SGLang image, install the same pin without dependency resolution
to avoid replacing its NumPy/CUDA stack:
```bash
pip install --no-deps "aiconfigurator==0.10.0"
```
Upgrade this pin only after rerunning the AIC predictor and compatibility tests.
## Quick start
The maintained tests define the supported first-version scope:
- [`test/test_simulation_sglang_runner.py`](test/test_simulation_sglang_runner.py):
direct Python use of the repository-level
[`SGLangBenchmarkRunner`](../../benchmark/simulator/bench_runner.py);
- [`test/test_simulation_sglang_serving.py`](test/test_simulation_sglang_serving.py):
server plus benchmark-client use through the HTTP serving path with AIC, ML,
and replay predictors and ShareGPT or timestamped traffic;
- [`test/test_simulation_offline_blocking.py`](test/test_simulation_offline_blocking.py):
equivalent logical results in `OFFLINE` and `BLOCKING` modes;
- [`test/test_simulation_cache_hit_ratio.py`](test/test_simulation_cache_hit_ratio.py):
reusable-prefix accounting and cache-tier hit metrics across repeated runs.
From `tools/sglang-simulator`:
```bash
python3 -m pytest -q test/test_simulation_sglang_runner.py
python3 -m pytest -q test/test_simulation_sglang_serving.py
```
Read these tests as the minimal maintained examples for constructing a dataset,
running a benchmark, starting a simulator server, sending programmatic, ShareGPT,
or timestamped traffic, comparing execution modes, and collecting request,
latency, throughput, and prefix-cache metrics.
## Serving mode
Choose a fresh output directory and export it in the server terminal before
starting the server:
```bash
export SGLANG_USE_CPU_ENGINE=1
export CUDA_VISIBLE_DEVICES=""
export SGLANG_SIMULATOR_OUTPUT_MODE=OFFLINE
export SIMULATOR_OUTPUT_DIR=/tmp/sglang-simulator-serving-001
test ! -e "$SIMULATOR_OUTPUT_DIR"
export SGLANG_SIMULATOR_OUTPUT_DIR="$SIMULATOR_OUTPUT_DIR"
python3 -m sglang_simulator.simulation.sglang.launch_server \
--model-path /absolute/path/to/model \
--sim-config-path /absolute/path/to/simulator.json \
--port 30000
```
In the benchmark terminal, export the same output directory before sending
timestamped traffic with the simulator-aware benchmark adapter:
```bash
cd /path/to/sglang
export SIMULATOR_OUTPUT_DIR=/tmp/sglang-simulator-serving-001
export SGLANG_SIMULATOR_OUTPUT_DIR="$SIMULATOR_OUTPUT_DIR"
python3 benchmark/simulator/bench_serving.py \
--simulator-mode offline \
--backend sglang \
--base-url http://127.0.0.1:30000 \
--model /absolute/path/to/model \
--dataset-name autobench \
--dataset-path /absolute/path/to/trace.jsonl \
--use-trace-timestamps \
--num-prompts 100 \
--profile \
--output-file "$SIMULATOR_OUTPUT_DIR/benchmark.json"
```
The server and benchmark are separate processes, so exporting
`SGLANG_SIMULATOR_OUTPUT_DIR` in the server terminal does not configure the
benchmark terminal. The benchmark adapter reads `metrics.json` from this path
after profiling and uses those server-side logical-time metrics for its serving
table and output file. If the benchmark points at another directory, it may show
unrelated stale metrics or client wall-clock values. Use the same fresh path in
both terminals for every run.
The simulator always runs the SGLang runtime with `tp_size=ep_size=dp_size=pp_size=1`
and both attention/decode context-parallel sizes set to `1`. Parallel CLI options
accepted by SGLang are therefore ignored by this simulator entry point. This keeps
simulator-only CPU work single-process; it does not change the modeled deployment.
Set the real deployment topology under `scheduler` in `--sim-config-path`. That
topology drives predictor and cache-resource modeling without launching physical
parallel workers.
Other server options are normal SGLang command-line arguments. For direct Python
integration, see
[`test_simulation_sglang_runner.py`](test/test_simulation_sglang_runner.py); for the
process/HTTP path, see
[`test_simulation_sglang_serving.py`](test/test_simulation_sglang_serving.py).
## Simulation modes
| Mode | Behavior |
|---|---|
| `OFFLINE` | Advances the simulator's logical clock without sleeping. |
| `BLOCKING` | Sleeps for predicted forward and visible L2-to-L1 load latency. |
Use server-side simulator metrics for comparisons. Client wall-clock duration is
not the simulated timeline in `OFFLINE` mode. When using the benchmark adapter,
make sure its `SGLANG_SIMULATOR_OUTPUT_DIR` matches the server's output directory
so the printed table and `benchmark.json` are sourced from the current run's
`metrics.json`.
## Configuration
A simulator configuration has three sections:
```json
{
"platform": {
"accelerator": {"name": "h20_sxm"},
"disk_read_bandwidth_gb": 8,
"disk_write_bandwidth_gb": 8,
"memory_read_bandwidth_gb": 64,
"memory_write_bandwidth_gb": 64,
"num_device_per_node": 1
},
"predictor": {
"name": "replay",
"database_path": "/absolute/path/to/replay_table.json"
},
"scheduler": {
"tp_size": 4,
"ep_size": 4,
"dp_size": 1,
"pp_size": 1,
"cp_size": 1,
"cp_style": "none",
"data_type": "BF16",
"kv_cache_data_type": "BF16",
"backend_name": "sglang"
}
}
```
- `platform` describes the simulated accelerator and storage bandwidth.
- `predictor` selects forward-latency prediction.
- `scheduler` describes the real target deployment topology and backend metadata.
`tp_size`, `ep_size`, `dp_size`, `pp_size`, and `cp_size` are modeled values;
they do not launch physical workers. For AIConfigurator, `tp_size` is converted
to attention TP after removing modeled DP and CP, while `cp_size` is passed as
AIConfigurator context parallelism. `cp_style` uses the AIConfigurator values
such as `none`, `allgather`, `ulysses`, or `ring`. Decode-only `dcp_size` has no
separate AIConfigurator field and is not modeled yet.
### Prefix-cache accuracy
Prefix-cache hit accuracy is highly sensitive to `max_total_tokens`. It controls
the simulated device KV-cache capacity and participates in hierarchical host-cache
sizing, so a mismatch changes eviction timing and device, host, and storage hit
attribution. For deployment-faithful results, copy `max_total_num_tokens=N` from
the real SGLang server startup log and launch the simulator with
`--max-total-tokens N`. Avoid relying on a separately estimated capacity when
comparing the simulator with production traces.
Supported predictors:
| Predictor | Purpose |
|---|---|
| `aiconfigurator` | Operator and module performance-database estimation. |
| `ml` | A trained sklearn-compatible 18-feature latency model. |
| `replay` | Exact or nearest-neighbor batch-composition replay. |
Relative predictor paths are resolved from the simulator configuration location.
Environment variables in paths use `${NAME}` syntax.
## Workload formats
The Autobench trace format uses timestamps in milliseconds:
```json
{"prompt":[1,2,3],"prompt_len":3,"output_len":1,"timestamp":200}
```
Random and ShareGPT workloads are also supported by the runner API and serving
benchmark paths.
## Validation
Run the CPU compatibility and unit tests from the repository root:
```bash
pip install -e tools/sglang-simulator
python3 -m pytest -q tools/sglang-simulator/test/test_simulation_sglang_runner.py
python3 -m pytest -q tools/sglang-simulator/test/test_simulation_sglang_serving.py
```
Run the two files as separate pytest commands because the runner test installs
process-global simulator hooks and state.
Run repository checks before submitting:
```bash
git ls-files -z tools/sglang-simulator | \
xargs -0 env SKIP=no-commit-to-branch pre-commit run --files
```
Runtime changes should also be validated in a matching official SGLang image
with both `OFFLINE` and `BLOCKING` modes. Predictor changes should report
step-level error, and scheduler or cache changes should compare request latency,
throughput, and prefix-cache reuse against measured traces.
+50
View File
@@ -0,0 +1,50 @@
# SGLang Simulator examples
The example assets are organized by purpose:
- `sim_configs/`: standalone AIC SOL, AIC SILICON, ML, and replay simulator configs;
- `assets/`: the small illustrative ML model, replay table, and test tokenizer;
- `workloads/`: ShareGPT and timestamped simulator/Autobench workload examples;
The ML model is an illustrative constant-latency sklearn model, not a calibrated
hardware predictor. Rebuild it and the tokenizer with:
```bash
python3 examples/build_example_assets.py
```
Only load pickle/joblib assets from sources you trust.
For maintained direct-run and serving examples, see
[`test_simulation_sglang_runner.py`](../test/test_simulation_sglang_runner.py) and
[`test_simulation_sglang_serving.py`](../test/test_simulation_sglang_serving.py).
Start a server with any example config:
```bash
python3 -m sglang_simulator.simulation.sglang.launch_server \
--model-path /path/to/model \
--sim-config-path examples/sim_configs/aic_sol.json \
--port 30000
```
Run a ShareGPT workload with at least four output tokens so decode and TPOT are
measured:
```bash
cd /path/to/sglang
python3 benchmark/simulator/bench_serving.py \
--simulator-mode=offline \
--backend=sglang \
--base-url=http://127.0.0.1:30000 \
--model=/path/to/model \
--tokenizer=/path/to/model \
--dataset-name=sharegpt \
--dataset-path=examples/workloads/sharegpt-example.json \
--sharegpt-output-len=4 \
--num-prompts=3 \
--profile
```
The timestamp trace uses the simulator-owned Autobench JSONL contract. Its
`timestamp` values are request-arrival times in milliseconds.
Binary file not shown.
@@ -0,0 +1,5 @@
{
"[[1, 3]]": 0.001,
"[[4, 0]]": 0.005,
"[[8, 0]]": 0.008
}
@@ -0,0 +1,59 @@
{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [
{
"id": 0,
"content": "[UNK]",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": false,
"special": true
}
],
"normalizer": null,
"pre_tokenizer": {
"type": "Whitespace"
},
"post_processor": {
"type": "TemplateProcessing",
"single": [
{
"Sequence": {
"id": "A",
"type_id": 0
}
}
],
"pair": [
{
"Sequence": {
"id": "A",
"type_id": 0
}
},
{
"Sequence": {
"id": "B",
"type_id": 1
}
}
],
"special_tokens": {}
},
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": {
"[UNK]": 0,
"prefix": 1,
"caching": 2,
"latency": 3,
"decode": 4,
"token": 5
},
"unk_token": "[UNK]"
}
}
@@ -0,0 +1,6 @@
{
"backend": "tokenizers",
"model_max_length": 1000000000000000019884624838656,
"tokenizer_class": "TokenizersBackend",
"unk_token": "[UNK]"
}
+55
View File
@@ -0,0 +1,55 @@
#!/usr/bin/env python3
"""Rebuild the small ML and tokenizer assets used by examples and tests."""
from pathlib import Path
import joblib
import numpy as np
from sglang_simulator.time_predictor.ml import MLTimePredictor
from sklearn.dummy import DummyRegressor
from tokenizers import Tokenizer
from tokenizers.models import WordLevel
from tokenizers.pre_tokenizers import Whitespace
from transformers import PreTrainedTokenizerFast
ASSETS = Path(__file__).parent / "assets"
def build_ml_model() -> None:
model = DummyRegressor(strategy="constant", constant=0.001)
model.fit(np.zeros((1, len(MLTimePredictor.FEATURE_NAMES))), [0.001])
joblib.dump(
{"model": model, "features": MLTimePredictor.FEATURE_NAMES},
ASSETS / "model.pkl",
)
def build_tokenizer() -> None:
tokenizer = Tokenizer(
WordLevel(
{
"[UNK]": 0,
"prefix": 1,
"caching": 2,
"latency": 3,
"decode": 4,
"token": 5,
},
unk_token="[UNK]",
)
)
tokenizer.pre_tokenizer = Whitespace()
PreTrainedTokenizerFast(
tokenizer_object=tokenizer,
unk_token="[UNK]",
).save_pretrained(ASSETS / "tokenizer")
def main() -> None:
ASSETS.mkdir(parents=True, exist_ok=True)
build_ml_model()
build_tokenizer()
if __name__ == "__main__":
main()
@@ -0,0 +1,21 @@
{
"platform": {
"accelerator": {"name": "a100_sxm", "hbm_capacity_gb": 80},
"disk_read_bandwidth_gb": 8,
"disk_write_bandwidth_gb": 8,
"memory_read_bandwidth_gb": 64,
"memory_write_bandwidth_gb": 64,
"num_device_per_node": 8
},
"predictor": {
"name": "aiconfigurator",
"database_mode": "SILICON"
},
"scheduler": {
"tp_size": 1,
"ep_size": 1,
"dp_size": 1,
"backend_name": "sglang",
"backend_version": "0.5.9"
}
}
@@ -0,0 +1,21 @@
{
"platform": {
"accelerator": {"name": "a100_sxm", "hbm_capacity_gb": 80},
"disk_read_bandwidth_gb": 8,
"disk_write_bandwidth_gb": 8,
"memory_read_bandwidth_gb": 64,
"memory_write_bandwidth_gb": 64,
"num_device_per_node": 8
},
"predictor": {
"name": "aiconfigurator",
"database_mode": "SOL"
},
"scheduler": {
"tp_size": 1,
"ep_size": 1,
"dp_size": 1,
"backend_name": "sglang",
"backend_version": "0.5.9"
}
}
@@ -0,0 +1,22 @@
{
"platform": {
"accelerator": {"name": "a100_sxm", "hbm_capacity_gb": 80},
"disk_read_bandwidth_gb": 8,
"disk_write_bandwidth_gb": 8,
"memory_read_bandwidth_gb": 64,
"memory_write_bandwidth_gb": 64,
"num_device_per_node": 8
},
"predictor": {
"name": "ml",
"database_path": "../assets/model.pkl",
"latency_scale": 1.0
},
"scheduler": {
"tp_size": 1,
"ep_size": 1,
"dp_size": 1,
"backend_name": "sglang",
"backend_version": "0.5.9"
}
}
@@ -0,0 +1,23 @@
{
"platform": {
"accelerator": {"name": "a100_sxm", "hbm_capacity_gb": 80},
"disk_read_bandwidth_gb": 8,
"disk_write_bandwidth_gb": 8,
"memory_read_bandwidth_gb": 64,
"memory_write_bandwidth_gb": 64,
"num_device_per_node": 8
},
"predictor": {
"name": "replay",
"database_path": "../assets/replay_table.json",
"miss_strategy": "knn",
"miss_knn_k": 1
},
"scheduler": {
"tp_size": 1,
"ep_size": 1,
"dp_size": 1,
"backend_name": "sglang",
"backend_version": "0.5.9"
}
}
@@ -0,0 +1,38 @@
[
{
"conversations": [
{
"from": "human",
"value": "Explain prefix caching in one concise sentence."
},
{
"from": "gpt",
"value": "Prefix caching reuses KV states shared by prompt prefixes."
}
]
},
{
"conversations": [
{
"from": "human",
"value": "What does time to first token measure?"
},
{
"from": "gpt",
"value": "It measures latency from request arrival to the first generated token."
}
]
},
{
"conversations": [
{
"from": "human",
"value": "Why does decode latency matter for serving?"
},
{
"from": "gpt",
"value": "Decode latency determines the cadence at which later tokens reach the user."
}
]
}
]
@@ -0,0 +1,3 @@
{"prompt":[100,101,102,103],"prompt_len":4,"output_len":4,"timestamp":0}
{"prompt":[100,101,102,104],"prompt_len":4,"output_len":4,"timestamp":75}
{"prompt":[200,201,202,203],"prompt_len":4,"output_len":4,"timestamp":250}
+18
View File
@@ -0,0 +1,18 @@
[build-system]
requires = ["setuptools>=42", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "sglang-simulator"
dynamic = ["version"]
description = "A simulation benchmark tool for sglang"
dependencies = [
"numpy",
"scikit-learn",
"joblib",
]
[project.optional-dependencies]
aic = [
"aiconfigurator==0.10.0"
]
+21
View File
@@ -0,0 +1,21 @@
from setuptools import find_packages, setup
def get_version():
version = "0.1.0"
with open("src/sglang_simulator/__init__.py") as f:
for line in f:
if line.startswith("__version__"):
version = line.split("=")[1].strip(' \n"')
return version
setup(
name="sglang-simulator",
version=get_version(),
url="https://github.com/sgl-project/sglang.git",
description="A High-Fidelity LLM inference simulator for SGLang",
packages=find_packages(where="src"),
package_dir={"": "src"},
py_modules=["usercustomize"],
)
@@ -0,0 +1 @@
__version__ = "0.1.0"
@@ -0,0 +1,99 @@
"""Early compatibility checks for the SGLang surfaces used by the simulator."""
import inspect
from importlib import metadata
SIMULATOR_SERVER_ARG_OVERRIDES = {
# The simulator models the target deployment topology separately through
# sim_config.scheduler. Keep the SGLang runtime single-process so host-side
# simulator work is not multiplied by the modeled parallel world size.
"tp_size": 1,
"ep_size": 1,
"dp_size": 1,
"pp_size": 1,
"attn_cp_size": 1,
"dcp_size": 1,
"disable_overlap_schedule": True,
"disable_cuda_graph": True,
"attention_backend": "torch_native",
"prefill_attention_backend": "torch_native",
"decode_attention_backend": "torch_native",
}
class SGLangCompatibilityError(RuntimeError):
pass
def _sglang_version() -> str:
try:
return metadata.version("sglang")
except metadata.PackageNotFoundError:
return "source-checkout"
def _require_parameters(function, required: set[str], surface: str) -> None:
parameters = set(inspect.signature(function).parameters)
missing = required - parameters
if missing:
raise SGLangCompatibilityError(
f"SGLang {_sglang_version()} is missing {surface} parameters: "
f"{', '.join(sorted(missing))}. The simulator must be adapted to "
"this SGLang revision before it can run."
)
def apply_simulator_server_args(target) -> None:
"""Apply simulator-owned values before constructing the final ServerArgs."""
if isinstance(target, dict):
target.update(SIMULATOR_SERVER_ARG_OVERRIDES)
return
for name, value in SIMULATOR_SERVER_ARG_OVERRIDES.items():
setattr(target, name, value)
def validate_simulator_server_args(server_args) -> None:
"""Fail early if a process bypassed a simulator-owned launch entry point."""
mismatches = [
f"{name}={getattr(server_args, name, None)!r} (expected {expected!r})"
for name, expected in SIMULATOR_SERVER_ARG_OVERRIDES.items()
if getattr(server_args, name, None) != expected
]
if mismatches:
raise SGLangCompatibilityError(
"SGLang Simulator server arguments were not prepared by a supported "
f"entry point: {', '.join(mismatches)}"
)
def validate_launch_runtime() -> None:
from sglang.srt.entrypoints.http_server import launch_server
_require_parameters(
launch_server,
{"server_args", "run_scheduler_process_func", "run_detokenizer_process_func"},
"launch_server",
)
def validate_benchmark_runtime() -> None:
from sglang.benchmark import serving
missing = [
name
for name in (
"BenchmarkMetrics",
"calculate_metrics",
"cli_main",
"get_request",
"run_benchmark",
)
if not hasattr(serving, name)
]
if missing:
raise SGLangCompatibilityError(
f"SGLang {_sglang_version()} is missing benchmark surfaces: "
f"{', '.join(missing)}. The simulator must be adapted to this "
"SGLang revision before it can run."
)
@@ -0,0 +1,35 @@
from sglang_simulator.dataset.base_dataset import (
BaseDataset,
GenericRequest,
SimpleDataset,
)
from sglang_simulator.dataset.dataset_args import DatasetArgs
from sglang_simulator.dataset.random import RandomDataset, RandomIDsDataset
from transformers import PreTrainedTokenizer
dataset_registry: dict[str, BaseDataset] = {
"random": RandomDataset,
"random_ids": RandomIDsDataset,
}
def get_dataset(
dataset_args: DatasetArgs, tokenizer: PreTrainedTokenizer | None = None
) -> BaseDataset:
if dataset_args.name not in dataset_registry:
raise ValueError(f"unknown dataset name: {dataset_args.name}")
dataset: BaseDataset = dataset_registry[dataset_args.name](
args=dataset_args, tokenizer=tokenizer
)
return dataset
__all__ = (
"DatasetArgs",
"BaseDataset",
"SimpleDataset",
"GenericRequest",
"get_dataset",
)
@@ -0,0 +1,242 @@
"""Simulator-owned loader for timestamped Autobench JSONL traces.
The trace format is a public SGLang Simulator input contract. Keep its parser
here instead of importing SGLang's benchmark-internal Autobench module, which
may be moved or removed independently of the simulator.
"""
import json
from argparse import Namespace
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Optional
import numpy as np
from transformers import PreTrainedTokenizerBase
from sglang.benchmark.datasets.common import BaseDataset, DatasetRow
_RESERVED_FIELDS = {
"prompt",
"messages",
"prompt_origin",
"output_len",
"max_tokens",
"max_completion_tokens",
"completion_tokens",
"prompt_len",
"text_prompt_len",
"vision_prompt_len",
"image_data",
"timestamp",
"routing_key",
"metadata",
"extra_request_body",
"param_send",
}
def _load_json_if_needed(value: Any) -> Any:
if not isinstance(value, str):
return value
value = value.strip()
if not value or value[0] not in "[{":
return value
try:
return json.loads(value)
except json.JSONDecodeError:
return value
def _normalize_messages(messages: Any) -> Optional[list[dict[str, Any]]]:
messages = _load_json_if_needed(messages)
if not isinstance(messages, list) or not messages:
return None
if not all(isinstance(message, dict) for message in messages):
return None
normalized = []
for message in messages:
if "role" not in message or message.get("content") is None:
return None
normalized.append({"role": message["role"], "content": message["content"]})
return normalized
def _normalize_prompt(row: dict[str, Any]) -> tuple[Any, str]:
for key in ("messages", "prompt_origin"):
normalized = _normalize_messages(row.get(key))
if normalized is not None:
return normalized, "messages"
prompt = _load_json_if_needed(row.get("prompt"))
if isinstance(prompt, list) and prompt:
if isinstance(prompt[0], dict):
normalized = _normalize_messages(prompt)
if normalized is not None:
return normalized, "messages"
if all(isinstance(item, int) for item in prompt):
return prompt, "token_ids"
if all(isinstance(item, str) for item in prompt):
return prompt, "multi_turn"
if all(
isinstance(turn, list)
and turn
and all(
isinstance(message, dict) and "role" in message and "content" in message
for message in turn
)
for turn in prompt
):
return prompt, "multi_turn"
if isinstance(prompt, str) and prompt:
return prompt, "prompt"
if isinstance(row.get("content"), list):
turns = [str(item) for item in row["content"]]
if len(turns) % 2 == 0:
turns = turns[:-1]
messages = []
if row.get("system"):
messages.append({"role": "system", "content": str(row["system"])})
messages.extend(
{
"role": "user" if index % 2 == 0 else "assistant",
"content": turn,
}
for index, turn in enumerate(turns)
)
if messages:
return messages, "messages"
raise ValueError("Unsupported Autobench row: missing prompt/messages")
def _prompt_lengths(
row: dict[str, Any],
prompt: Any,
prompt_kind: str,
tokenizer: Optional[PreTrainedTokenizerBase],
) -> tuple[int, int, int]:
if row.get("prompt_len") is not None:
prompt_len = int(row["prompt_len"])
return (
prompt_len,
int(row.get("text_prompt_len", prompt_len)),
int(row.get("vision_prompt_len", 0)),
)
if prompt_kind == "token_ids":
return len(prompt), len(prompt), 0
if tokenizer is None:
raise ValueError("Autobench rows without prompt_len require a tokenizer")
if prompt_kind == "messages":
prompt_len = len(
tokenizer.apply_chat_template(
prompt, tokenize=True, add_generation_prompt=True
)
)
return prompt_len, prompt_len, 0
if prompt_kind == "prompt":
prompt_len = len(tokenizer.encode(prompt, add_special_tokens=False))
return prompt_len, prompt_len, 0
return 0, 0, 0
def _extra_request_body(row: dict[str, Any]) -> dict[str, Any]:
extra = {}
param_send = _load_json_if_needed(row.get("param_send"))
if isinstance(param_send, dict):
extra.update(param_send)
extra.update(
{key: value for key, value in row.items() if key not in _RESERVED_FIELDS}
)
explicit = _load_json_if_needed(row.get("extra_request_body"))
if isinstance(explicit, dict):
extra.update(explicit)
return extra
def sample_autobench_requests(
dataset_path: str,
num_requests: int,
tokenizer: Optional[PreTrainedTokenizerBase],
fixed_output_len: Optional[int] = None,
) -> list[DatasetRow]:
dataset = []
with Path(dataset_path).open(encoding="utf-8") as file:
for line_number, line in enumerate(file, start=1):
if num_requests > 0 and len(dataset) >= num_requests:
break
if not line.strip():
continue
try:
row = json.loads(line)
prompt, prompt_kind = _normalize_prompt(row)
prompt_len, text_prompt_len, vision_prompt_len = _prompt_lengths(
row, prompt, prompt_kind, tokenizer
)
except (TypeError, ValueError, json.JSONDecodeError) as error:
raise ValueError(
f"Invalid Autobench row {line_number} in {dataset_path}: {error}"
) from error
output_len = fixed_output_len
for key in (
"output_len",
"max_tokens",
"max_completion_tokens",
"completion_tokens",
):
output_len = output_len or row.get(key)
dataset.append(
DatasetRow(
prompt=prompt,
prompt_len=prompt_len,
output_len=int(output_len or 256),
text_prompt_len=text_prompt_len,
vision_prompt_len=vision_prompt_len,
image_data=row.get("image_data"),
timestamp=row.get("timestamp"),
routing_key=row.get("routing_key"),
extra_request_body=_extra_request_body(row),
)
)
print(f"Loaded {len(dataset)} Autobench requests")
print(f"#Input tokens: {np.sum([row.prompt_len for row in dataset])}")
print(f"#Output tokens: {np.sum([row.output_len for row in dataset])}")
return dataset
@dataclass
class AutoBenchmarkDataset(BaseDataset):
dataset_path: str
num_requests: int
fixed_output_len: Optional[int]
@classmethod
def from_args(cls, args: Namespace) -> "AutoBenchmarkDataset":
return cls(
dataset_path=args.dataset_path,
num_requests=args.num_prompts,
fixed_output_len=getattr(args, "sharegpt_output_len", None),
)
def load(
self,
tokenizer: PreTrainedTokenizerBase,
model_id: Optional[str] = None,
) -> list[DatasetRow]:
return sample_autobench_requests(
self.dataset_path,
self.num_requests,
tokenizer,
self.fixed_output_len,
)
def register_autobench_dataset() -> None:
"""Register the simulator trace contract with SGLang's serving benchmark."""
from sglang.benchmark import datasets
datasets.DATASET_MAPPING["autobench"] = AutoBenchmarkDataset
@@ -0,0 +1,72 @@
from dataclasses import dataclass, field
from typing import Optional, overload
from sglang_simulator.dataset.dataset_args import DatasetArgs
from transformers import PreTrainedTokenizerBase
@dataclass(slots=True)
class GenericRequest:
prompt: Optional[str] = None
token_ids: Optional[list[int]] = None
input_length: int = -1
output_length: int = -1
custom_params: dict = field(default_factory=dict)
def __post_init__(self):
if self.prompt is None and self.token_ids is None:
raise ValueError("Invalid Request")
class BaseDataset:
def __init__(self, tokenizer: PreTrainedTokenizerBase, args: DatasetArgs):
self.tokenizer: PreTrainedTokenizerBase = tokenizer
self.args = args
self._name = ""
@overload
def __getitem__(self, index: int) -> GenericRequest: ...
@overload
def __getitem__(self, index: slice) -> list[GenericRequest]: ...
def __getitem__(self, index):
"""Get item(s) by index or slice. Delegates to _get_single_item for single items."""
if isinstance(index, slice):
start, stop, step = index.indices(len(self))
return [self[i] for i in range(start, stop, step)]
if index >= len(self):
raise IndexError
return self._get_single_item(index)
def _get_single_item(self, index: int) -> GenericRequest:
raise NotImplementedError
def __len__(self) -> int:
raise NotImplementedError
@property
def name(self):
if self._name:
return self._name
else:
return self.__class__.__name__
class SimpleDataset(BaseDataset):
def __init__(
self, tokenizer=None, args=None, reqs: list[GenericRequest] | None = None
):
super().__init__(tokenizer, args)
self.data: list[GenericRequest] = []
if reqs is not None:
self.data.extend(reqs)
def add_request(self, req: GenericRequest):
self.data.append(req)
def _get_single_item(self, index: int) -> GenericRequest:
return self.data[index]
def __len__(self):
return len(self.data)
@@ -0,0 +1,20 @@
from dataclasses import dataclass
@dataclass
class DatasetArgs:
name: str = ""
filepath: str = ""
num_prompts: int = -1
min_input_len: int = -1
max_input_len: int = -1
min_output_len: int = -1
max_output_len: int = -1
@property
def mean_input_length(self):
return (self.min_input_len + self.max_input_len) // 2
@property
def mean_output_length(self):
return (self.min_output_len + self.max_output_len) // 2
@@ -0,0 +1,52 @@
from random import randint
from sglang_simulator.dataset.base_dataset import BaseDataset, GenericRequest
class RandomIDsDataset(BaseDataset):
def __init__(self, tokenizer, args):
super().__init__(tokenizer, args)
self.cached: list[GenericRequest] = []
self._name = "random_ids"
def __len__(self):
return self.args.num_prompts
def _get_single_item(self, index: int) -> GenericRequest:
if index < len(self.cached):
return self.cached[index]
min_id, max_id = (
int(self.tokenizer.vocab_size * 0.25),
int(self.tokenizer.vocab_size * 0.75),
)
input_len = randint(self.args.min_input_len, self.args.max_input_len)
input_ids = [randint(min_id, max_id) for _ in range(input_len)]
req = GenericRequest(
token_ids=input_ids,
input_length=input_len,
output_length=randint(self.args.min_output_len, self.args.max_output_len),
)
self.cached.append(req)
return req
class RandomDataset(RandomIDsDataset):
def __init__(self, tokenizer, args):
super().__init__(tokenizer, args)
self.cached: list[GenericRequest] = []
self._name = "random"
def __len__(self):
return self.args.num_prompts
def _get_single_item(self, index: int) -> GenericRequest:
if index < len(self.cached):
return self.cached[index]
req = super()._get_single_item(index)
if req.token_ids is not None:
req.prompt = self.tokenizer.decode(req.token_ids, skip_special_tokens=True)
req.token_ids = None
self.cached.append(req)
return req
@@ -0,0 +1,15 @@
from sglang_simulator.hook.base_hook import BaseHook
from sglang_simulator.hook.class_hook_entry import (
install_class_hooks,
is_class_hook_matched,
remove_class_hooks,
validate_required_class_hooks,
)
__all__ = (
install_class_hooks,
is_class_hook_matched,
remove_class_hooks,
validate_required_class_hooks,
BaseHook,
)
@@ -0,0 +1,36 @@
from typing import List, Optional, Union
from sglang_simulator.utils import get_logger
logger = get_logger("sgl_simulator")
class BaseHook:
HOOK_CLASS_NAME: Optional[str] = None
HOOK_MODULE_NAME: Optional[str] = None
REGEX: bool = False
REQUIRED: bool = True
def __init__(self):
pass
@classmethod
def hook(cls, target) -> None:
"""
Return a new target or simply modify the target reference.
"""
raise NotImplementedError
def _register_hooks(HOOKS: List[BaseHook], hooks: Union[List[BaseHook], BaseHook]):
if isinstance(hooks, list):
for hook in hooks:
if not issubclass(hook, BaseHook):
raise TypeError("The hook should inherit from BaseHook.")
HOOKS.append(hook)
elif isinstance(hooks, type) and issubclass(hooks, BaseHook):
HOOKS.append(hooks)
else:
raise TypeError(
"The type of registered hook should be a list of BaseHook or a single BaseHook."
)
@@ -0,0 +1,80 @@
import builtins
import re
from types import FunctionType
from typing import List, Union
from sglang_simulator.hook.base_hook import BaseHook, _register_hooks
from sglang_simulator.utils import get_logger
logger = get_logger("sgl_simulator")
CLASS_HOOKS: List[BaseHook] = []
_MATCHED_CLASS_HOOKS = set()
_builtins_build_class_ = builtins.__build_class__
def _custom_build_class_(func, name: str, *bases, **kwargs):
for hook in CLASS_HOOKS:
if (
hook.REGEX
and hook.HOOK_CLASS_NAME
and re.search(hook.HOOK_CLASS_NAME, name)
) or name == hook.HOOK_CLASS_NAME:
module_name = None
if isinstance(func, FunctionType):
module_name = getattr(func, "__globals__", {}).get("__name__", "")
if (
hook.REGEX and re.search(hook.HOOK_MODULE_NAME, module_name)
) or module_name == hook.HOOK_MODULE_NAME:
logger.debug(
f"Hooking Class: {hook.__name__} into {module_name}|{name}"
+ (
"(Regex is enabled, which might cause unexpected behavior.)"
if hook.REGEX
else ""
)
)
target_class = _builtins_build_class_(func, name, *bases, **kwargs)
hook.hook(target_class)
_MATCHED_CLASS_HOOKS.add(hook)
return target_class
return _builtins_build_class_(func, name, *bases, **kwargs)
def install_class_hooks(hooks: Union[List[BaseHook], BaseHook]):
_register_hooks(CLASS_HOOKS, hooks)
builtins.__build_class__ = _custom_build_class_
def is_class_hook_matched(hook: BaseHook) -> bool:
return hook in _MATCHED_CLASS_HOOKS
def validate_required_class_hooks() -> None:
unmatched = [
hook
for hook in CLASS_HOOKS
if hook.REQUIRED and hook not in _MATCHED_CLASS_HOOKS
]
if not unmatched:
return
hook_names = ", ".join(
f"{hook.__name__} ({hook.HOOK_MODULE_NAME}.{hook.HOOK_CLASS_NAME})"
for hook in unmatched
)
raise RuntimeError(
"Required SGLang Simulator hooks did not match imported SGLang classes: "
f"{hook_names}. The simulator must be adapted to this SGLang revision."
)
def remove_class_hooks():
# Clear the registered hooks and reset the build class function.
# Note: The classes that have been hooked will not be reset.
CLASS_HOOKS.clear()
_MATCHED_CLASS_HOOKS.clear()
builtins.__build_class__ = _builtins_build_class_
@@ -0,0 +1,8 @@
def get_obj_from_args(type_name: str, *args, **kwargs):
for obj in args:
if type_name == f"{type(obj).__module__}.{type(obj).__name__}":
return obj
for obj in kwargs.values():
if type_name == f"{type(obj).__module__}.{type(obj).__name__}":
return obj
return None
@@ -0,0 +1,4 @@
from sglang_simulator.simulation.benchmark.base_runner import BaseBenchmarkRunner
from sglang_simulator.simulation.benchmark.bench_config import BenchmarkConfig
__all__ = ["BaseBenchmarkRunner", "BenchmarkConfig"]
@@ -0,0 +1,18 @@
from abc import ABC, abstractmethod
class BaseBenchmarkRunner(ABC):
def __init__(self):
pass
@abstractmethod
def benchmark(self) -> dict:
pass
@abstractmethod
def flush_cache(self):
pass
@abstractmethod
def shutdown(self):
pass
@@ -0,0 +1,9 @@
from dataclasses import dataclass
from typing import Optional
@dataclass
class BenchmarkConfig:
request_rate: float = float("inf")
max_concurrency: Optional[int] = None
ignore_request_timestamp: bool = False
@@ -0,0 +1,5 @@
from sglang_simulator.simulation.manager.config import ConfigManager
from sglang_simulator.simulation.manager.env import Envs
from sglang_simulator.simulation.manager.state import StateManager
__all__ = ["StateManager", "Envs", "ConfigManager"]
@@ -0,0 +1,252 @@
import json
from pathlib import Path
from typing import Optional
from sglang_simulator.simulation.manager.env import Envs
from sglang_simulator.simulation.types import PlatformConfig, SchedulerConfig
from sglang_simulator.simulation.utils import (
calc_kv_cache_cell_elems,
calc_kv_cache_per_layer_elems,
)
from sglang_simulator.spec import AcceleratorInfo, DataType, ModelInfo
from sglang_simulator.time_predictor import (
AIConfiguratorTimePredictor,
InferTimePredictor,
)
from sglang_simulator.utils import get_logger
logger = get_logger()
class ConfigManager:
"""Centralized configuration manager with caching."""
_model_info: Optional[ModelInfo] = None
_platform_config: Optional[PlatformConfig] = None
_scheduler_config: Optional[SchedulerConfig] = None
_raw_config: Optional[dict] = None
@classmethod
def _get_raw_config(cls) -> dict:
if cls._raw_config is None:
with open(Envs.config_path()) as f:
cls._raw_config = json.load(f)
return cls._raw_config
@classmethod
def resolve_config_relative_path(cls, path: str | None) -> str | None:
"""Resolve predictor assets without depending on the process cwd."""
if not path or Path(path).is_absolute():
return path
cwd_candidate = Path(path)
if cwd_candidate.exists():
return str(cwd_candidate.resolve())
config_path = Path(Envs.config_path()).resolve()
for parent in config_path.parents:
candidate = parent / path
if candidate.exists():
return str(candidate)
# Keep the original value so predictor-specific errors remain clear.
return path
@classmethod
def reset_config_cache(cls):
cls._raw_config = None
cls._model_info = None
cls._platform_config = None
cls._scheduler_config = None
@classmethod
def set_model_info(cls, model: ModelInfo):
cls._model_info = model
@classmethod
def get_model_info(cls) -> ModelInfo | None:
return cls._model_info
@classmethod
def get_accelerator_info(cls) -> AcceleratorInfo:
config = cls._get_raw_config()
platform_config = config.get("platform", {})
acc_info = platform_config.get("accelerator", {})
hw = AcceleratorInfo.find_by_hw_name(acc_info.get("name"))
if hw is None:
logger.debug(
f"Failed to initialize device info with {acc_info.get('name')}. All available devices are: {AcceleratorInfo.list_all_hws().keys()}"
)
hw = AcceleratorInfo(
name=acc_info.get("name"),
vendor=acc_info.get("vendor"),
hbm_bandwidth_gb=acc_info.get("hbm_bandwidth_gb"),
hbm_capacity_gb=acc_info.get("hbm_capacity_gb"),
inter_node_bandwidth_gb=acc_info.get("inter_node_bandwidth_gb"),
intra_node_bandwidth_gb=acc_info.get("intra_node_bandwidth_gb"),
tflops=acc_info.get("tflops"),
)
else:
logger.info(f"Device info initialized: {hw}")
return hw
@classmethod
def get_platform_config(cls) -> PlatformConfig:
if cls._platform_config is None:
hw = cls.get_accelerator_info()
config = cls._get_raw_config()
platform_config = config.get("platform", {})
cls._platform_config = PlatformConfig(
device=hw,
disk_read_bandwidth_gb=platform_config.get("disk_read_bandwidth_gb"),
disk_write_bandwidth_gb=platform_config.get("disk_write_bandwidth_gb"),
memory_read_bandwidth_gb=platform_config.get(
"memory_read_bandwidth_gb"
),
memory_write_bandwidth_gb=platform_config.get(
"memory_write_bandwidth_gb"
),
num_device_per_node=platform_config.get("num_device_per_node"),
)
logger.info(
f"Platform configuration initialized successfully. {cls._platform_config}"
)
return cls._platform_config
@classmethod
def set_scheduler_config(cls, config: SchedulerConfig):
# The configuration from the external config file has higher priority.
external_config = cls._get_raw_config().get("scheduler", {})
for field_name in [
"tp_size",
"dp_size",
"ep_size",
"pp_size",
"cp_size",
"cp_style",
"backend_name",
"backend_version",
"kv_bytes_per_token_per_gpu",
"hicache_ratio",
"moe_quant_mode_override",
"fmha_quant_mode_override",
"comm_quant_mode_override",
]:
field_value = external_config.get(field_name)
if field_value is not None:
setattr(config, field_name, field_value)
for field_name in ["data_type", "kv_cache_data_type"]:
field_value = external_config.get(field_name)
if field_value is not None:
setattr(config, field_name, DataType(field_value))
cls._scheduler_config = config
@classmethod
def get_kv_cache_bytes(cls) -> int:
model = cls._model_info
scheduler_config = cls._scheduler_config
return (
calc_kv_cache_cell_elems(
model, scheduler_config.tp_size, scheduler_config.pp_size
)
* scheduler_config.kv_cache_data_type.bytes
)
@classmethod
def get_kv_cache_bytes_per_layer(cls) -> int:
model = cls._model_info
scheduler_config = cls._scheduler_config
return (
calc_kv_cache_per_layer_elems(
model, scheduler_config.tp_size, scheduler_config.pp_size
)
* scheduler_config.kv_cache_data_type.bytes
)
@classmethod
def get_scheduler_config(cls):
return cls._scheduler_config
@classmethod
def _parse_server_args(cls, server_args: dict, backend: str) -> SchedulerConfig:
if backend == "sglang":
return SchedulerConfig(
tp_size=server_args.get("tp_size", 1),
ep_size=server_args.get("ep_size", 1),
dp_size=server_args.get("dp_size", 1),
pp_size=server_args.get("pp_size", 1),
cp_size=server_args.get("attn_cp_size", 1),
cp_style=server_args.get("cp_style", "none"),
mem_fraction_static=server_args.get("mem_fraction_static"),
backend_name="sglang",
)
else:
raise RuntimeError(f"Unsupported backend[{backend}] server args parser.")
@classmethod
def get_inference_time_predictor(
cls, model: ModelInfo, hw: AcceleratorInfo, sched_config: SchedulerConfig
) -> InferTimePredictor:
config = cls._get_raw_config()
predictor_config = config.get("predictor", {})
if predictor_config.get("name") == "aiconfigurator":
database_mode = predictor_config.get("database_mode", "SILICON")
prefill_scale_factor = predictor_config.get("prefill_scale_factor", 1)
decode_scale_factor = predictor_config.get("decode_scale_factor", 1)
prefill_min_latency = predictor_config.get("prefill_min_latency", 0)
workload_distribution = predictor_config.get(
"workload_distribution", "balanced"
)
enable_oom_check = predictor_config.get("enable_oom_check", False)
return AIConfiguratorTimePredictor(
model,
hw=hw,
config=sched_config,
database_path=cls.resolve_config_relative_path(
predictor_config.get("database_path")
),
database_mode=database_mode,
prefill_scale_factor=prefill_scale_factor,
decode_scale_factor=decode_scale_factor,
prefill_min_latency=prefill_min_latency,
workload_distribution=workload_distribution,
enable_oom_check=enable_oom_check,
)
elif predictor_config.get("name") == "ml":
from sglang_simulator.time_predictor.ml import MLTimePredictor
return MLTimePredictor(
model,
hw=hw,
config=sched_config,
database_path=cls.resolve_config_relative_path(
predictor_config.get("database_path")
),
latency_scale=predictor_config.get("latency_scale", 1.0),
)
elif predictor_config.get("name") == "replay":
from sglang_simulator.time_predictor.replay import ReplayTimePredictor
return ReplayTimePredictor(
model,
hw=hw,
config=sched_config,
database_path=cls.resolve_config_relative_path(
predictor_config.get("database_path")
),
miss_fallback_seconds=predictor_config.get(
"miss_fallback_seconds", 0.0
),
miss_strategy=predictor_config.get("miss_strategy", "zero"),
miss_knn_k=predictor_config.get("miss_knn_k", 3),
)
else:
raise ValueError(
f"Unknown predictor name: {predictor_config.get('name')}. "
f"Supported: aiconfigurator, ml, replay"
)
@@ -0,0 +1,52 @@
import os
from sglang_simulator.utils.logger import get_logger
logger = get_logger("sgl_simulator")
class Envs:
@classmethod
def config_path(cls) -> str:
SGLANG_SIMULATOR_CONFIG_PATH = os.getenv("SGLANG_SIMULATOR_CONFIG_PATH")
if not SGLANG_SIMULATOR_CONFIG_PATH or not os.path.exists(
SGLANG_SIMULATOR_CONFIG_PATH
):
raise RuntimeError(
f"The mock configuration path is not set or does not exist({SGLANG_SIMULATOR_CONFIG_PATH}). Please set it using the system variable SGLANG_SIMULATOR_CONFIG_PATH"
)
return SGLANG_SIMULATOR_CONFIG_PATH
@classmethod
def output_dir(cls) -> str:
SGLANG_SIMULATOR_OUTPUT_DIR = os.getenv(
"SGLANG_SIMULATOR_OUTPUT_DIR", "/tmp/sglang_simulator/output/"
)
SGLANG_SIMULATOR_OUTPUT_DIR = os.path.realpath(SGLANG_SIMULATOR_OUTPUT_DIR)
if os.path.exists(SGLANG_SIMULATOR_OUTPUT_DIR) and os.path.isfile(
SGLANG_SIMULATOR_OUTPUT_DIR
):
logger.error(
f"The metrics output path, {SGLANG_SIMULATOR_OUTPUT_DIR}, exists and is a file."
)
raise RuntimeError(
f"{SGLANG_SIMULATOR_OUTPUT_DIR} exists but is not a directory."
)
os.makedirs(SGLANG_SIMULATOR_OUTPUT_DIR, exist_ok=True)
return SGLANG_SIMULATOR_OUTPUT_DIR
@classmethod
def hicache_storage_keys_path(cls) -> str:
SGLANG_SIMULATOR_HICACHE_STORAGE_KEYS_PATH = os.getenv(
"SGLANG_SIMULATOR_HICACHE_STORAGE_KEYS_PATH",
"/tmp/sglang_simulator/hicache_storage_keys.txt",
)
return SGLANG_SIMULATOR_HICACHE_STORAGE_KEYS_PATH
@classmethod
def simulation_mode(cls) -> str:
SGLANG_SIMULATOR_OUTPUT_MODE = os.getenv(
"SGLANG_SIMULATOR_OUTPUT_MODE", "OFFLINE"
).upper()
assert SGLANG_SIMULATOR_OUTPUT_MODE in ("BLOCKING", "OFFLINE")
return SGLANG_SIMULATOR_OUTPUT_MODE
@@ -0,0 +1,123 @@
class StateManager:
_iteration: int = 0
_global_clock: float = 0
_last_inference_dur: float = 0
_current_inference_dur: float = 0
_hicache_l2_load_dur: float = 0
_hicache_l2_backup_dur: float = 0
_hicache_l2_load_call_count: int = 0
_hicache_l2_load_segment_count: int = 0
_hicache_l2_load_units: int = 0
_hicache_l2_load_bytes: float = 0
_last_real_time_ts: float = 0
_last_flush_time_ts: float = 0
@classmethod
def reset(cls):
cls._iteration = 0
cls._global_clock = 0
cls._last_inference_dur = 0
cls._current_inference_dur = 0
cls._hicache_l2_backup_dur = 0
cls._hicache_l2_load_dur = 0
cls._hicache_l2_load_call_count = 0
cls._hicache_l2_load_segment_count = 0
cls._hicache_l2_load_units = 0
cls._hicache_l2_load_bytes = 0
cls._last_real_time_ts = 0
@classmethod
def inc_iteration(cls) -> None:
cls._iteration += 1
@classmethod
def get_iteration(cls) -> int:
return cls._iteration
@classmethod
def inc_hicache_l2_load_dur(cls, dur: float) -> None:
cls._hicache_l2_load_dur += dur
@classmethod
def inc_hicache_l2_load_stats(
cls,
call_count: int = 0,
segment_count: int = 0,
units: int = 0,
bytes_: float = 0,
) -> None:
cls._hicache_l2_load_call_count += call_count
cls._hicache_l2_load_segment_count += segment_count
cls._hicache_l2_load_units += units
cls._hicache_l2_load_bytes += bytes_
@classmethod
def inc_hicache_l2_backup_dur(cls, dur: float) -> None:
cls._hicache_l2_backup_dur += dur
@classmethod
def pop_hicache_l2_load_dur(cls) -> float:
dur = cls._hicache_l2_load_dur
cls._hicache_l2_load_dur = 0
return dur
@classmethod
def pop_hicache_l2_load_stats(cls) -> dict:
stats = {
"h2d_load_call_count": cls._hicache_l2_load_call_count,
"h2d_load_segment_count": cls._hicache_l2_load_segment_count,
"h2d_load_units": cls._hicache_l2_load_units,
"h2d_load_bytes": cls._hicache_l2_load_bytes,
}
cls._hicache_l2_load_call_count = 0
cls._hicache_l2_load_segment_count = 0
cls._hicache_l2_load_units = 0
cls._hicache_l2_load_bytes = 0
return stats
@classmethod
def pop_hicache_l2_backup_dur(cls) -> float:
dur = cls._hicache_l2_backup_dur
cls._hicache_l2_backup_dur = 0
return dur
@classmethod
def get_global_clock(cls) -> float:
return cls._global_clock
@classmethod
def step_global_clock(cls, dur: float) -> None:
cls._global_clock += dur
@classmethod
def set_global_clock(cls, clock: float) -> None:
cls._global_clock = clock
@classmethod
def set_current_inference_dur(cls, dur: float) -> None:
cls._last_inference_dur = cls._current_inference_dur
cls._current_inference_dur = dur
@classmethod
def get_last_inference_dur(cls) -> float:
return cls._last_inference_dur
@classmethod
def get_current_inference_dur(cls) -> float:
return cls._current_inference_dur
@classmethod
def set_last_real_time_ts(cls, ts):
cls._last_real_time_ts = ts
@classmethod
def get_last_real_time_ts(cls):
return cls._last_real_time_ts
@classmethod
def set_last_flush_time_ts(cls, ts: float):
cls._last_flush_time_ts = ts
@classmethod
def get_last_flush_time_ts(cls) -> float:
return cls._last_flush_time_ts
@@ -0,0 +1,311 @@
from queue import Empty, Queue
from typing import Optional
from sglang_simulator.hook import BaseHook
from sglang_simulator.simulation.manager import ConfigManager, StateManager
from sglang_simulator.simulation.sglang.req_stats_manager import request_stats_manager
class C_HiCacheController(BaseHook):
HOOK_CLASS_NAME = "HiCacheController"
HOOK_MODULE_NAME = "sglang.srt.managers.cache_controller"
REQUIRED = False
KV_CACHE_BYTES: Optional[int] = None
DISK_READ_BANDWIDTH_BYTES: Optional[float] = None
DISK_WRITE_BANDWIDTH_BYTES: Optional[float] = None
@staticmethod
def calc_prefetch_pages(
required_pages: int, page_size_byte: int, max_dur: float, bandwidth: float
) -> tuple[float, float]:
_prefetch_dur = required_pages * page_size_byte / bandwidth
if _prefetch_dur > max_dur:
_completed_pages = max(max_dur * bandwidth / page_size_byte, 1)
return _completed_pages, max_dur
else:
return required_pages, _prefetch_dur
@classmethod
def hook(cls, target):
original_terminate_prefetch = target.terminate_prefetch
original_storage_hit_query = target._storage_hit_query
original_init = target.__init__
original_append_host_mem_release = target.append_host_mem_release
def wrapped_init(self, *args, **kwargs):
self.sim_prefetch_buffer = Queue()
result = original_init(self, *args, **kwargs)
# The real IO thread normally creates this queue. The simulator
# replaces that thread, so initialize the handoff queue here.
if hasattr(self, "prefetch_hit_queue"):
self.prefetch_buffer = Queue()
return result
def wrapped_append_host_mem_release(self, host_indices):
# A terminated prefetch may not have allocated host memory yet.
if host_indices is None:
return
return original_append_host_mem_release(self, host_indices)
def override_backup_thread_func(self, *args, **kwargs):
# Async thread: perform no action
# The action will be performed by `handle_backup_operation`
pass
def override_prefetch_thread_func(self, *args, **kwargs):
# Async thread: perform no action
# The action will be performed by `handle_prefetch_operation`
pass
def handle_backup_operation(self):
if not self.enable_storage:
return
while True:
try:
operation = self.backup_queue.get(block=False)
if operation is None:
return
if not self.backup_skip:
self._page_backup(operation)
# TODO: Track the backup operation according to the global clock
self.ack_backup_queue.put(operation)
except Empty:
return
def handle_prefetch_operation(self):
if not self.enable_storage:
return
if C_HiCacheController.KV_CACHE_BYTES is None:
C_HiCacheController.KV_CACHE_BYTES = ConfigManager.get_kv_cache_bytes()
if C_HiCacheController.DISK_READ_BANDWIDTH_BYTES is None:
C_HiCacheController.DISK_READ_BANDWIDTH_BYTES = (
ConfigManager.get_platform_config().disk_read_bandwidth
)
# TODO: Overlap schedule
remain_dur = StateManager.get_current_inference_dur()
# Process all operations in the prefetch_queue: place those meeting
# the prefetch criteria into the sim_prefetch_buffer, and release the
# remaining operations along with any excess memory they have allocated.
while not self.prefetch_queue.empty():
try:
operation = self.prefetch_queue.get(block=False)
if operation is None:
break
# Ignore terminated operation
if operation._terminated_flag:
if hasattr(self, "prefetch_revoke_queue"):
self.prefetch_revoke_queue.put(operation.request_id)
else:
self.append_host_mem_release(operation.host_indices)
continue
hash_value, storage_hit_count = self._storage_hit_query(operation)
# not to prefetch if not enough benefits
if (
self.prefetch_threshold is not None
and storage_hit_count < self.prefetch_threshold
):
if hasattr(self, "prefetch_revoke_queue"):
self.prefetch_revoke_queue.put(operation.request_id)
continue
operation.mark_terminate()
self.append_host_mem_release(operation.host_indices)
continue
else:
operation.hash_value = hash_value[
: (storage_hit_count // self.page_size)
]
if hasattr(self, "prefetch_hit_queue"):
# Allocate only the storage-hit range on the scheduler.
operation.storage_hit_count = storage_hit_count
self.prefetch_hit_queue.put(operation)
continue
storage_hit_count = (
storage_hit_count // self.page_size * self.page_size
)
# free the pre-allocated memory for pages that are not hit
self.append_host_mem_release(
operation.host_indices[storage_hit_count:]
)
operation.host_indices = operation.host_indices[
:storage_hit_count
]
self.sim_prefetch_buffer.put(operation)
except Empty:
break
# handle operation which not yet fully prefetched
chunked_prefetch_operation = getattr(
self, "chunked_prefetch_operation", None
)
if chunked_prefetch_operation is not None:
operation = chunked_prefetch_operation["operation"]
if operation._terminated_flag:
setattr(self, "chunked_prefetch_operation", None)
self.append_host_mem_release(
operation.host_indices[int(operation.completed_tokens) :]
)
else:
storage_hit_count = chunked_prefetch_operation["storage_hit_count"]
completed_tokens, prefetch_dur = (
C_HiCacheController.calc_prefetch_pages(
(storage_hit_count - operation.completed_tokens),
C_HiCacheController.KV_CACHE_BYTES,
remain_dur,
C_HiCacheController.DISK_READ_BANDWIDTH_BYTES,
)
)
if (
completed_tokens
< storage_hit_count - operation.completed_tokens
):
operation.completed_tokens += completed_tokens
remain_dur = 0
else:
operation.completed_tokens = int(storage_hit_count)
operation.mark_terminate()
remain_dur -= prefetch_dur
setattr(self, "chunked_prefetch_operation", None)
# Release host memory after current operation is finished
self.append_host_mem_release(
operation.host_indices[storage_hit_count:]
)
# Feed operations whose host pages were allocated by the scheduler
# into the virtual-time transfer loop.
prefetch_buffer = getattr(self, "prefetch_buffer", None)
if prefetch_buffer is not None:
while not prefetch_buffer.empty():
try:
operation = prefetch_buffer.get(block=False)
if operation is not None:
self.sim_prefetch_buffer.put(operation)
except Empty:
break
# handle operation in sim_prefetch_buffer
while remain_dur > 0:
try:
operation = self.sim_prefetch_buffer.get(block=False)
if operation is None:
return
# Ignore terminated operation
if operation._terminated_flag:
self.append_host_mem_release(
operation.host_indices[int(operation.completed_tokens) :]
)
continue
storage_hit_count = len(operation.host_indices)
completed_tokens, prefetch_dur = (
C_HiCacheController.calc_prefetch_pages(
storage_hit_count,
C_HiCacheController.KV_CACHE_BYTES,
remain_dur,
C_HiCacheController.DISK_READ_BANDWIDTH_BYTES,
)
)
if completed_tokens < storage_hit_count:
# Continue to prefetch data next time.
operation.completed_tokens = completed_tokens
setattr(
self,
"chunked_prefetch_operation",
{
"operation": operation,
"storage_hit_count": storage_hit_count,
},
)
remain_dur = 0
else:
operation.completed_tokens = int(
storage_hit_count // self.page_size * self.page_size
)
# TODO: Track the prefetch operation according to the global clock
operation.mark_terminate()
remain_dur -= prefetch_dur
except Empty:
return
def override_generic_page_set(
self, hash_values, host_indices, extra_info=None
) -> bool:
host_pool = getattr(self, "storage_host_pool", self.mem_pool_host)
# Always pass extra_info to storage_backend.
data = [
host_pool.get_data_page(host_indices[i * self.page_size])
for i in range(len(hash_values))
]
return self.storage_backend.batch_set(hash_values, data, extra_info)
def wrapped_terminate_prefetch(self, operator):
result = original_terminate_prefetch(self, operator)
# This value may be a float if prefetch progress is interrupted by HiRadixCache.check_prefetch_progress.
result = (int(result[0]), result[1])
# operation.completed_tokens, operation.hash_value = result
req_stats = request_stats_manager.get_req_stats(operator.request_id)
req_stats.final_storage_hit_len = result[0]
return result
def wrapped_storage_hit_query(self, operator):
result = original_storage_hit_query(self, operator)
# hash_value, storage_query_count = result
req_stats = request_stats_manager.get_req_stats(operator.request_id)
req_stats.recv_storage_hit_len = result[1]
return result
target.__init__ = wrapped_init
target.prefetch_thread_func = override_prefetch_thread_func
target.backup_thread_func = override_backup_thread_func
target.handle_backup_operation = handle_backup_operation
target.handle_prefetch_operation = handle_prefetch_operation
target.append_host_mem_release = wrapped_append_host_mem_release
target._generic_page_set = override_generic_page_set
target.terminate_prefetch = wrapped_terminate_prefetch
target.storage_hit_query = wrapped_storage_hit_query
if hasattr(target, "_storage_hit_query"):
target._storage_hit_query = wrapped_storage_hit_query
class C_HybridCacheController(BaseHook):
"""Adapt UnifiedRadixCache's controller without duplicating legacy logic."""
HOOK_CLASS_NAME = "HybridCacheController"
HOOK_MODULE_NAME = "sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller"
REQUIRED = False
@classmethod
def hook(cls, target):
# HybridCacheController inherits the deterministic thread replacements and
# handle_* methods installed on HiCacheController. Its own initialization
# creates Unified's control queues after the base initializer returns.
original_init = target.__init__
original_storage_hit_query = target._storage_hit_query
def wrapped_init(self, *args, **kwargs):
result = original_init(self, *args, **kwargs)
if hasattr(self, "prefetch_hit_queue") and not hasattr(
self, "prefetch_buffer"
):
self.prefetch_buffer = Queue()
return result
def wrapped_storage_hit_query(self, operator):
result = original_storage_hit_query(self, operator)
req_stats = request_stats_manager.get_req_stats(operator.request_id)
req_stats.recv_storage_hit_len = result[1]
return result
target.__init__ = wrapped_init
target._storage_hit_query = wrapped_storage_hit_query
@@ -0,0 +1,19 @@
"""SGLang engine entry point with simulator-aware worker processes."""
from sglang_simulator.simulation.sglang.hook_bootstrap import (
install_simulator_hooks,
run_simulator_detokenizer_process,
run_simulator_scheduler_process,
)
install_simulator_hooks()
# Install hooks before importing Engine so its worker entry points are patched.
from sglang.srt.entrypoints.engine import Engine # noqa: E402
class SGLangSimulationEngine(Engine):
"""Engine whose spawned workers install SGLang Simulator hooks explicitly."""
run_scheduler_process_func = staticmethod(run_simulator_scheduler_process)
run_detokenizer_process_func = staticmethod(run_simulator_detokenizer_process)
@@ -0,0 +1,155 @@
import os
from typing import Any, List, Optional
from sglang_simulator.hook import BaseHook
from sglang_simulator.simulation.manager.env import Envs
from sglang_simulator.utils.logger import get_logger
logger = get_logger("sglang-simulator")
class C_StorageBackendFactory(BaseHook):
HOOK_CLASS_NAME = "StorageBackendFactory"
HOOK_MODULE_NAME = "sglang.srt.mem_cache.storage.backend_factory"
REQUIRED = False
@classmethod
def hook(cls, target):
def override_create_backend(cls, *args, **kwargs):
logger.info("Creating hijacked cache storage backend.")
return MockHiCacheStorage()
target.create_backend = override_create_backend
class MockHiCacheStorage:
def __init__(self, *args, **kwargs):
self.storage: set = set()
self.storage_file_path: str = Envs.hicache_storage_keys_path()
os.makedirs(os.path.dirname(self.storage_file_path), exist_ok=True)
if os.path.exists(self.storage_file_path):
with open(self.storage_file_path) as f:
line = f.readline()
while line:
self.storage.add(line.strip())
line = f.readline()
self.registered_pools = {}
def register_mem_pool_host(self, mem_pool_host):
pass
def register_mem_host_pool_v2(self, host_pool, host_pool_name):
"""Register one pool from UnifiedRadixCache's multi-pool HiCache stack."""
self.registered_pools[host_pool_name] = host_pool
@staticmethod
def _pool_storage_key(key: str, pool_name) -> str:
name = str(pool_name)
return key if name == "kv" else f"{key}.{name}"
def set(
self,
key: str,
value: Optional[Any] = None,
target_location: Optional[Any] = None,
target_sizes: Optional[Any] = None,
) -> bool:
if self.exists(key):
return True
self.storage.add(key)
with open(self.storage_file_path, "a+") as f:
f.write(key + "\n")
return True
def batch_set(
self,
keys: List[str],
values: Optional[Any] = None,
extra_info=None, # HiCacheStorageExtraInfo
target_locations: Optional[Any] = None,
target_sizes: Optional[Any] = None,
) -> bool:
for key, value in zip(keys, values):
if not self.set(key, value):
return False
return True
def exists(self, key: str) -> bool:
return key in self.storage
def batch_exists(self, keys: List[str], extra_info) -> int:
for i in range(len(keys)):
if not self.exists(keys[i]):
return i
return len(keys)
def batch_exists_v2(self, keys, pool_transfers=None, extra_info=None):
"""Return Unified HiCache's per-pool longest-prefix result."""
from sglang.srt.mem_cache.hicache_storage import PoolTransferResult
kv_hit_pages = self.batch_exists(keys, extra_info)
extra_pool_hit_pages = {}
final_pages = kv_hit_pages
for transfer in pool_transfers or []:
def has_component(page_idx):
return self.exists(
self._pool_storage_key(keys[page_idx], transfer.name)
)
hit_policy = getattr(transfer.hit_policy, "value", transfer.hit_policy)
if hit_policy == "all_pages":
boundary = next(
(i for i in range(kv_hit_pages) if not has_component(i)),
kv_hit_pages,
)
else:
trailing = max(1, len(transfer.keys) if transfer.keys else 1)
boundary = 0
for prefix_len in range(kv_hit_pages, 0, -1):
if all(
has_component(i)
for i in range(max(0, prefix_len - trailing), prefix_len)
):
boundary = prefix_len
break
extra_pool_hit_pages[transfer.name] = boundary
final_pages = min(final_pages, boundary)
return PoolTransferResult(
kv_hit_pages=final_pages,
extra_pool_hit_pages=extra_pool_hit_pages,
)
def batch_get_v2(self, transfers, extra_info=None):
"""Simulate loading every available pool page into registered host pools."""
results = {}
for transfer in transfers:
keys = transfer.keys or []
results[transfer.name] = [
self.exists(self._pool_storage_key(key, transfer.name)) for key in keys
]
return results
def batch_set_v2(self, transfers, extra_info=None):
"""Persist Unified HiCache component keys without materializing payloads."""
results = {}
for transfer in transfers:
keys = transfer.keys or []
pool_results = []
for key in keys:
pool_results.append(
self.set(self._pool_storage_key(key, transfer.name))
)
results[transfer.name] = pool_results
return results
def clear(self) -> bool:
self.storage.clear()
with open(self.storage_file_path, "w"):
pass
return True
@@ -0,0 +1,25 @@
from sglang_simulator.hook import BaseHook
class C_HiRadixCacheHook(BaseHook):
HOOK_CLASS_NAME = "HiRadixCache"
HOOK_MODULE_NAME = "sglang.srt.mem_cache.hiradix_cache"
REQUIRED = False
@classmethod
def hook(cls, target):
original_check_hicache_events = target.check_hicache_events
def wrapped_check_hicache_events(self, *args, **kwargs):
# The async thread for prefetching and backup in `HiCacheController` has been deprecated.
# So we have to handle the backup or prefetch operation manually.
self.cache_controller.handle_backup_operation()
self.cache_controller.handle_prefetch_operation()
result = original_check_hicache_events(self, *args, **kwargs)
# Host pages are allocated after the storage query. Run the
# simulated transfer only after that allocation step.
if hasattr(self.cache_controller, "prefetch_hit_queue"):
self.cache_controller.handle_prefetch_operation()
return result
target.check_hicache_events = wrapped_check_hicache_events
@@ -0,0 +1,80 @@
"""Install SGLang Simulator hooks in the parent and spawned SGLang worker processes."""
import os
# Spawned interpreters inherit this marker before usercustomize runs.
os.environ["SGLANG_SIMULATOR_BOOTSTRAP"] = "1"
import sglang_simulator.hook as sglang_simulator_hook
from sglang_simulator.simulation.sglang import (
cache_controller,
hicache_storage,
hiradix_cache,
mem_cache_allocator,
mem_pool_host,
model_runner,
scheduler,
sgl_kernel_hook,
unified_radix_cache,
)
# A spawned worker imports this module while unpickling its target. ModelConfig
# can import GPU kernels while later arguments are still being unpickled, before
# the target wrapper executes, so the loader stub must already be present here.
sgl_kernel_hook.install_load_utils_stub()
_HOOKS_INSTALLED = False
def install_simulator_hooks() -> None:
"""Install hooks once in the current Python interpreter."""
global _HOOKS_INSTALLED
if _HOOKS_INSTALLED:
return
# The package __init__ loads GPU ops before a child-module import hook can
# run reliably under spawn. Seed the loader module before importing SGLang.
sgl_kernel_hook.install_load_utils_stub()
sglang_simulator_hook.install_class_hooks(
[
scheduler.C_SchedulerHook,
scheduler.C_SglangPrefillAdderHook,
scheduler.C_SchedulerRequestReceiver,
model_runner.C_ModelRunnerHook,
model_runner.C_KVCacheConfiguratorHook,
hicache_storage.C_StorageBackendFactory,
cache_controller.C_HiCacheController,
cache_controller.C_HybridCacheController,
hiradix_cache.C_HiRadixCacheHook,
unified_radix_cache.C_UnifiedRadixCacheHook,
mem_cache_allocator.C_PagedTokenToKVPoolAllocatorHook,
mem_pool_host.C_MHATokenToKVPoolHostHook,
mem_pool_host.C_HostKVCacheHook,
mem_pool_host.C_PackedSingleKVPoolHook,
mem_pool_host.C_GenericHostKVCacheSubclassHook,
]
)
_HOOKS_INSTALLED = True
def run_simulator_scheduler_process(*args, **kwargs):
"""Spawn-safe scheduler entry point which installs SGLang Simulator before SGLang imports."""
install_simulator_hooks()
# Spawned workers do not inherit parent-process monkey patches, so import
# the scheduler only after installing hooks in this process.
from sglang.srt.managers.scheduler import run_scheduler_process
return run_scheduler_process(*args, **kwargs)
def run_simulator_detokenizer_process(*args, **kwargs):
"""Spawn-safe detokenizer entry point which installs SGLang Simulator before imports."""
install_simulator_hooks()
# Install hooks before transitive schedule_batch and memory_pool imports
# so this CPU-only process does not load real GPU kernels.
from sglang.srt.managers.detokenizer_manager import run_detokenizer_process
return run_detokenizer_process(*args, **kwargs)
@@ -0,0 +1,109 @@
import argparse
import dataclasses
import os
import sys
from typing import Optional
from sglang_simulator.compat import (
apply_simulator_server_args,
validate_launch_runtime,
)
from sglang_simulator.simulation.sglang.hook_bootstrap import (
install_simulator_hooks,
run_simulator_detokenizer_process,
run_simulator_scheduler_process,
)
from sglang_simulator.utils import get_logger
install_simulator_hooks()
logger = get_logger("sgl_simulator")
@dataclasses.dataclass
class SimulationArgs:
sim_config_path: Optional[str] = None
@staticmethod
def add_cli_args(parser: argparse.ArgumentParser):
parser.add_argument(
"--sim-config-path",
type=str,
default=None,
help="Path to simulation JSON config (same as SGLANG_SIMULATOR_CONFIG_PATH).",
)
@classmethod
def from_cli_args(cls, ns: argparse.Namespace) -> "SimulationArgs":
return SimulationArgs(sim_config_path=ns.sim_config_path)
def _has_cli_option(argv: list[str], option: str) -> bool:
return any(arg == option or arg.startswith(f"{option}=") for arg in argv)
def apply_simulator_defaults(raw_args: argparse.Namespace, argv: list[str]) -> None:
"""Avoid real model execution while preserving explicit SGLang options."""
if not _has_cli_option(argv, "--load-format"):
raw_args.load_format = "dummy"
if os.getenv("SGLANG_USE_CPU_ENGINE") != "1":
return
if not _has_cli_option(argv, "--device"):
raw_args.device = "cpu"
if not _has_cli_option(argv, "--attention-backend"):
raw_args.attention_backend = "torch_native"
if not _has_cli_option(argv, "--sampling-backend"):
raw_args.sampling_backend = "pytorch"
if not (
_has_cli_option(argv, "--cuda-graph-backend-decode")
or _has_cli_option(argv, "--cuda-graph-backend-prefill")
or _has_cli_option(argv, "--disable-cuda-graph")
):
raw_args.disable_cuda_graph = True
# CPU-only model validation may still query CUDA capability while
# constructing ServerArgs, before the simulator runner is spawned.
import torch
torch.cuda.get_device_capability = lambda *_args, **_kwargs: (10, 0)
if __name__ == "__main__":
validate_launch_runtime()
from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import kill_process_tree
parser = argparse.ArgumentParser()
g = parser.add_argument_group("sglang")
ServerArgs.add_cli_args(g)
g = parser.add_argument_group("simulation")
SimulationArgs.add_cli_args(g)
argv = sys.argv[1:]
raw_args = parser.parse_args(argv)
apply_simulator_defaults(raw_args, argv)
apply_simulator_server_args(raw_args)
server_args = ServerArgs.from_cli_args(raw_args)
simulation_args = SimulationArgs.from_cli_args(raw_args)
config_path = os.getenv("SGLANG_SIMULATOR_CONFIG_PATH")
if config_path and os.path.exists(config_path):
logger.info(f"Using config from {config_path}")
elif simulation_args.sim_config_path:
os.environ["SGLANG_SIMULATOR_CONFIG_PATH"] = simulation_args.sim_config_path
try:
launch_server(
server_args,
run_scheduler_process_func=run_simulator_scheduler_process,
run_detokenizer_process_func=run_simulator_detokenizer_process,
)
finally:
kill_process_tree(os.getpid(), include_parent=False)
@@ -0,0 +1,128 @@
import types
import torch
from sglang_simulator.hook import BaseHook
def _alloc_extend_cpu(
self,
prefix_lens: torch.Tensor,
prefix_lens_cpu: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_cpu: torch.Tensor,
last_loc: torch.Tensor,
extend_num_tokens: int,
num_new_pages: int = None,
):
"""CPU implementation using SGLang's native paged-allocation helper."""
from sglang.srt.mem_cache.allocator import alloc_extend_naive
from sglang.srt.utils import get_num_new_pages
if num_new_pages is None:
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
prefix_lens=prefix_lens_cpu,
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
if num_new_pages > len(self.free_pages):
return None
out_indices = torch.empty(
(extend_num_tokens,),
dtype=self.free_pages.dtype,
device=self.device,
)
alloc_extend_naive(
prefix_lens,
seq_lens,
last_loc,
self.free_pages,
out_indices,
self.page_size,
self.device,
)
self.free_pages = self.free_pages[num_new_pages:]
return out_indices
def _alloc_decode_cpu(
self,
seq_lens: torch.Tensor,
seq_lens_cpu: torch.Tensor,
last_loc: torch.Tensor,
):
"""CPU decode allocation through the allocator's public method contract."""
from sglang.srt.utils import get_num_new_pages
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
decode=True,
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
if num_new_pages > len(self.free_pages):
return None
out_indices = (last_loc + 1).to(dtype=self.free_pages.dtype)
need_new_page = seq_lens % self.page_size == 1
if num_new_pages:
out_indices = out_indices.clone()
out_indices[need_new_page] = self.free_pages[:num_new_pages] * self.page_size
self.free_pages = self.free_pages[num_new_pages:]
return out_indices
def alloc_extend_cpu(*args, **kwargs):
"""Compatibility entry plus the native allocator-method implementation."""
if args and isinstance(args[0], torch.Tensor):
from sglang.srt.mem_cache.allocator import alloc_extend_naive
prefix_lens, seq_lens, last_loc, free_pages, out_indices = args[:5]
alloc_extend_naive(
prefix_lens,
seq_lens,
last_loc,
free_pages,
out_indices,
kwargs["page_size"],
prefix_lens.device,
)
return None
return _alloc_extend_cpu(*args, **kwargs)
def alloc_decode_cpu(*args, **kwargs):
"""Compatibility entry plus the native allocator-method implementation."""
if args and isinstance(args[0], torch.Tensor):
seq_lens, last_loc, free_pages, out_indices = args[:4]
page_size = kwargs["page_size"]
need_new_page = seq_lens % page_size == 1
result = last_loc + 1
result[need_new_page] = (
free_pages[: int(need_new_page.sum().item())] * page_size
)
out_indices.copy_(result)
return None
return _alloc_decode_cpu(*args, **kwargs)
class C_PagedTokenToKVPoolAllocatorHook(BaseHook):
HOOK_CLASS_NAME = "PagedTokenToKVPoolAllocator"
HOOK_MODULE_NAME = r"^sglang\.srt\.mem_cache\.allocator(?:\.paged)?$"
REGEX = True
@classmethod
def hook(cls, target):
original_init = target.__init__
def wrapped_init(self, *args, **kwargs):
original_init(self, *args, **kwargs)
if self.device == "cpu":
self.alloc_extend = types.MethodType(_alloc_extend_cpu, self)
self.alloc_decode = types.MethodType(_alloc_decode_cpu, self)
target.__init__ = wrapped_init
@@ -0,0 +1,384 @@
from abc import ABC, abstractmethod
from enum import Enum
from functools import lru_cache
import numpy as np
import torch
from sglang_simulator.hook import BaseHook
from sglang_simulator.simulation.manager import ConfigManager, StateManager
from sglang_simulator.utils import get_logger
logger = get_logger()
class TransportDirection(Enum):
H2D = "H2D"
D2H = "D2H"
class HicacheTransportEstimator(ABC):
def __init__(
self,
memory_read_bandwidth_bytes: float,
memory_write_bandwidth_bytes: float,
):
self.memory_read_bandwidth_bytes = memory_read_bandwidth_bytes
self.memory_write_bandwidth_bytes = memory_write_bandwidth_bytes
@abstractmethod
def estimate_bandwidth(
self, size_bytes: np.ndarray, direction: TransportDirection
) -> np.ndarray:
raise NotImplementedError
class HicacheTransportOverheadEstimator(HicacheTransportEstimator):
"""Bandwidth model with a fixed launch overhead and 85% efficiency."""
def estimate_bandwidth(
self, size_bytes: np.ndarray, direction: TransportDirection
) -> np.ndarray:
if direction is TransportDirection.H2D:
overhead_s = 6.67e-6
bandwidth = self.memory_read_bandwidth_bytes * 0.85
else:
overhead_s = 4e-6
bandwidth = self.memory_write_bandwidth_bytes * 0.85
return size_bytes * bandwidth / (overhead_s * bandwidth + size_bytes)
def compute_contiguous_index_lengths(
host_indices: torch.Tensor,
device_indices: torch.Tensor,
) -> np.ndarray:
if len(host_indices) != len(device_indices):
raise ValueError("Host and device cache index lists must have the same length.")
if len(host_indices) == 0:
return np.empty(0, dtype=np.float64)
host = np.asarray(host_indices.cpu(), dtype=np.int64)
device = np.asarray(device_indices.cpu(), dtype=np.int64)
contiguous = (np.diff(host) == 1) & (np.diff(device) == 1)
cuts = np.flatnonzero(~contiguous) + 1
starts = np.r_[0, cuts]
ends = np.r_[cuts, len(host_indices)]
return (ends - starts).astype(np.float64)
def allocate_meta_tensor(
dims,
dtype: torch.dtype,
device: str,
pin_memory: bool,
allocator=None,
registration_granularity_bytes=None,
) -> torch.Tensor:
"""Allocate metadata-only host cache payload for simulation."""
return torch.empty(dims, dtype=dtype, device="meta")
def _install_meta_allocators() -> None:
modules = []
try:
from sglang.srt.mem_cache import memory_pool_host
modules.append(memory_pool_host)
except ImportError:
pass
try:
from sglang.srt.mem_cache.pool_host import common
modules.append(common)
except ImportError:
pass
for module in modules:
allocators = getattr(module, "ALLOC_MEMORY_FUNCS", None)
if allocators is None:
continue
allocators.default_factory = lambda: allocate_meta_tensor
for key in list(allocators):
allocators[key] = allocate_meta_tensor
_SIMULATED_AVAILABLE_HOST_MEMORY_BYTES = 1 << 60
class _PsutilProxy:
def __init__(self, psutil_module):
self._psutil_module = psutil_module
def virtual_memory(self):
snapshot = self._psutil_module.virtual_memory()
return snapshot._replace(
available=max(
snapshot.available,
_SIMULATED_AVAILABLE_HOST_MEMORY_BYTES,
)
)
def __getattr__(self, name):
return getattr(self._psutil_module, name)
def _call_with_meta_host_memory(original_init, self, *args, **kwargs):
"""Bypass physical host-payload checks while meta allocation is active."""
init_globals = getattr(original_init, "__globals__", None)
psutil_module = init_globals.get("psutil") if init_globals is not None else None
if psutil_module is None:
return original_init(self, *args, **kwargs)
proxy = _PsutilProxy(psutil_module)
init_globals["psutil"] = proxy
try:
return original_init(self, *args, **kwargs)
finally:
if init_globals.get("psutil") is proxy:
init_globals["psutil"] = psutil_module
@lru_cache(maxsize=256)
def get_refined_cache_size_per_token(host_pool) -> float:
internal_size = float(host_pool.get_size_per_token())
scheduler_config = ConfigManager.get_scheduler_config()
if scheduler_config is None or scheduler_config.kv_cache_data_type is None:
logger.warning(
"Scheduler KV-cache dtype is unavailable; using %s's native "
"size-per-token value.",
host_pool.__class__.__name__,
)
return internal_size
internal_dtype = host_pool.dtype
dtype_factor = scheduler_config.kv_cache_data_type.bytes / internal_dtype.itemsize
return internal_size * dtype_factor
_DSV4_TRANSFER_SIZE_MULTIPLIERS = {
"swa": 130,
"deepseek_v4_c4": 65,
"deepseek_v4_c4_indexer": 132,
"deepseek_v4_c128": 3,
"deepseek_v4_c4_state": 256,
"deepseek_v4_c128_state": 256,
"deepseek_v4_indexer_state": 128,
"deepseek_v4_c4_indexer_state": 128,
}
_DSV4_PAGED_POOL_NAMES = {
"swa",
"deepseek_v4_c4",
"deepseek_v4_c4_indexer",
"deepseek_v4_c128",
}
def _dsv4_transfer_size_multiplier(host_pool) -> int | None:
return _DSV4_TRANSFER_SIZE_MULTIPLIERS.get(str(getattr(host_pool, "pool_name", "")))
def get_transfer_size_per_unit(host_pool, *, all_layers: bool) -> float:
"""Return calibrated bytes moved for one transfer unit.
DSv4's paged and state pools expose physical page-row geometry through Unified
HiCache. Preserve the 0714 estimator's calibrated logical-byte multipliers while
keeping transfer dispatch on the current Unified pool interfaces.
"""
size = get_refined_cache_size_per_token(host_pool)
dsv4_multiplier = _dsv4_transfer_size_multiplier(host_pool)
if dsv4_multiplier is not None:
return size * dsv4_multiplier
layer_num = max(int(getattr(host_pool, "layer_num", 1)), 1)
per_layer_size = size / layer_num
return per_layer_size * layer_num if all_layers else per_layer_size
def _transport_estimator() -> HicacheTransportEstimator:
platform = ConfigManager.get_platform_config()
return HicacheTransportOverheadEstimator(
memory_read_bandwidth_bytes=platform.memory_read_bandwidth,
memory_write_bandwidth_bytes=platform.memory_write_bandwidth,
)
def _normalize_transfer_indices(self, host_indices, device_indices):
if host_indices is None or device_indices is None:
return None, None
if hasattr(self, "_to_page_indices"):
host_indices = self._to_page_indices(host_indices)
device_indices = self._to_page_indices(device_indices)
return host_indices, device_indices
def _transfer_segment_lengths(
self, host_indices, device_indices, *, count_logical_tokens: bool = False
) -> np.ndarray:
if host_indices is None or device_indices is None:
return np.empty(0, dtype=np.float64)
original_unit_count = len(host_indices)
host_indices, device_indices = _normalize_transfer_indices(
self, host_indices, device_indices
)
lengths = compute_contiguous_index_lengths(host_indices, device_indices)
if (
len(lengths)
and count_logical_tokens
and str(getattr(self, "pool_name", "")) in _DSV4_PAGED_POOL_NAMES
):
# The 0714 DSv4 paged-pool H2D estimator counted logical token slots
# while using page-row contiguity to determine transfer segments.
lengths[-1] += original_unit_count - len(host_indices)
return lengths
def _sim_load_to_device_per_layer(
self,
device_pool,
host_indices,
device_indices,
layer_id,
io_backend,
*,
is_draft: bool = False,
) -> None:
segment_lengths = _transfer_segment_lengths(
self, host_indices, device_indices, count_logical_tokens=True
)
if not len(segment_lengths):
return
size_bytes = segment_lengths * get_transfer_size_per_unit(self, all_layers=False)
StateManager.inc_hicache_l2_load_stats(
call_count=1,
segment_count=len(size_bytes),
units=int(np.sum(segment_lengths)),
bytes_=float(np.sum(size_bytes)),
)
bandwidth = _transport_estimator().estimate_bandwidth(
size_bytes, TransportDirection.H2D
)
StateManager.inc_hicache_l2_load_dur(float(np.sum(size_bytes / bandwidth)))
def _sim_backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
) -> None:
segment_lengths = _transfer_segment_lengths(self, host_indices, device_indices)
if not len(segment_lengths):
return
size_bytes = segment_lengths * get_transfer_size_per_unit(self, all_layers=True)
bandwidth = _transport_estimator().estimate_bandwidth(
size_bytes, TransportDirection.D2H
)
StateManager.inc_hicache_l2_backup_dur(float(np.sum(size_bytes / bandwidth)))
def _sim_get_data_page(self, index, flat: bool = True) -> torch.Tensor:
return torch.ones(size=(1, 1)) * index
def _sim_set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None:
return None
def _install_transport_methods(target) -> None:
original_init = target.__init__
def wrapped_init(self, *args, **kwargs):
_install_meta_allocators()
if "pin_memory" in kwargs:
kwargs["pin_memory"] = False
return _call_with_meta_host_memory(original_init, self, *args, **kwargs)
target.__init__ = wrapped_init
target.load_to_device_per_layer = _sim_load_to_device_per_layer
target.backup_from_device_all_layer = _sim_backup_from_device_all_layer
target.get_data_page = _sim_get_data_page
target.set_from_flat_data_page = _sim_set_from_flat_data_page
class C_MHATokenToKVPoolHostHook(BaseHook):
HOOK_CLASS_NAME = "MHATokenToKVPoolHost"
HOOK_MODULE_NAME = r"^sglang\.srt\.mem_cache\.(memory_pool_host|pool_host\.mha)$"
REGEX = True
REQUIRED = False
@classmethod
def hook(cls, target):
_install_transport_methods(target)
class C_HostKVCacheHook(BaseHook):
HOOK_CLASS_NAME = "HostKVCache"
HOOK_MODULE_NAME = r"^sglang\.srt\.mem_cache\.(memory_pool_host|pool_host\.base)$"
REGEX = True
REQUIRED = False
@classmethod
def hook(cls, target):
original_init = target.__init__
def wrapped_init(self, *args, **kwargs):
_install_meta_allocators()
if "pin_memory" in kwargs:
kwargs["pin_memory"] = False
elif len(args) > 5:
args = list(args)
args[5] = False
return _call_with_meta_host_memory(original_init, self, *args, **kwargs)
target.__init__ = wrapped_init
class C_PackedSingleKVPoolHook(BaseHook):
"""Allocate byte-packed single KV pools from their runtime geometry."""
HOOK_CLASS_NAME = r".*SingleKVPool$"
HOOK_MODULE_NAME = r"^sglang\.srt\.mem_cache\..+$"
REGEX = True
REQUIRED = False
@classmethod
def hook(cls, target):
original_create_buffer = target.create_buffer
def wrapped_create_buffer(self, *, num_pages: int):
if self.store_dtype != torch.uint8 or not hasattr(
self, "get_bytes_per_token"
):
return original_create_buffer(self, num_pages=num_pages)
try:
return original_create_buffer(self, num_pages=num_pages)
except AssertionError:
# Some packed pools validate production-only geometry. Simulation
# still needs a correctly sized byte buffer for dummy model configs.
pass
bytes_per_token = self.get_bytes_per_token()
self.kv_cache_total_dim = bytes_per_token
bytes_per_page = self.page_size * bytes_per_token
self.bytes_per_page_padded = (bytes_per_page + 575) // 576 * 576
return torch.zeros(
num_pages,
self.bytes_per_page_padded,
dtype=self.store_dtype,
device=self.device,
)
target.create_buffer = wrapped_create_buffer
class C_GenericHostKVCacheSubclassHook(BaseHook):
HOOK_CLASS_NAME = r".*(?:PoolHost|HostPool)$"
HOOK_MODULE_NAME = r"^sglang\.srt\.mem_cache\.(memory_pool_host|pool_host\..+)$"
REGEX = True
REQUIRED = False
@classmethod
def hook(cls, target):
if any(base.__name__ == "HostKVCache" for base in target.__mro__[1:]):
_install_transport_methods(target)
@@ -0,0 +1,265 @@
import torch
from sglang_simulator.hook import BaseHook
from sglang_simulator.simulation.manager import ConfigManager
from sglang_simulator.simulation.sglang.utils import (
resolve_model_info,
resolve_scheduler_config,
)
from sglang_simulator.simulation.utils import profile_device_available_bytes
class _MockModel(torch.nn.Module):
"""Minimal model surface needed by SGLang's native runner initialization."""
def forward(self, *args, **kwargs):
return None
class _MockModelLoader:
"""Minimal loader state for upstream resident-weight accounting."""
preloaded_weights_bytes = 0
def _make_mock_model_loader(model_runner_type):
if hasattr(model_runner_type, "preloaded_weights_bytes"):
return _MockModelLoader()
return None
def _resolve_kv_page_size(configurator):
return (
getattr(configurator, "page_size", None)
or getattr(configurator.server_args, "page_size", None)
or 1
)
class C_ModelRunnerHook(BaseHook):
HOOK_CLASS_NAME = "ModelRunner"
HOOK_MODULE_NAME = "sglang.srt.model_executor.model_runner"
@classmethod
def hook(cls, target):
def override_load_model(self):
from sglang.srt.model_executor.model_runner import (
resolve_sliding_window_size,
)
self.model = _MockModel()
self.dtype = self.model_config.dtype
self.sliding_window_size = resolve_sliding_window_size(
self.model, self.model_config
)
self.prefill_aware_swa = False
self.weight_load_mem_usage = 0
self.load_config = None
self.loader = _make_mock_model_loader(type(self))
if ConfigManager.get_model_info() is None:
ConfigManager.set_model_info(resolve_model_info(self.model_config))
def wrapped_forward(self, *args, **kwargs):
batch = args[0]
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.model_executor.model_runner import ModelRunnerOutput
output = LogitsProcessorOutput(
next_token_logits=torch.empty(
size=(batch.batch_size, self.model_config.vocab_size),
device=self.device,
)
)
return ModelRunnerOutput(
logits_output=output,
can_run_graph=False,
expert_distribution_metrics=None,
)
def wrapped_sample(self, *args, **kwargs):
logits = args[0]
return torch.ones(
size=(logits.next_token_logits.shape[0],),
device=self.device,
dtype=torch.int64,
)
def wrapped_compute_logprobs_only(*args, **kwargs):
return None
def wrapped_init_attention_backends(self):
try:
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
resolve_attention_backend_strs,
)
except ImportError:
default_backend = self.server_args.attention_backend
self.prefill_attention_backend_str = (
self.server_args.prefill_attention_backend or default_backend
)
self.decode_attention_backend_str = (
self.server_args.decode_attention_backend or default_backend
)
else:
resolved = resolve_attention_backend_strs(model_runner=self)
self.prefill_attention_backend_str = resolved.prefill
self.decode_attention_backend_str = resolved.decode
self.attn_backend = None
self.decode_attn_backend = None
self.decode_attn_backend_group = None
def wrapped_init_cuda_graphs(self, capture_decode_cuda_graph=True):
self.graph_mem_usage = 0
self.cuda_graph_runner = None
self.eager_runner = None
self.prefill_cuda_graph_runner = None
self.decode_cuda_graph_runner = None
# Keep SGLang's native initialize() and alloc_memory_pool() lifecycle.
# Only operations that require model weights or GPU kernels are mocked.
target.load_model = override_load_model
target.forward = wrapped_forward
target.sample = wrapped_sample
target.compute_logprobs_only = wrapped_compute_logprobs_only
target.init_attention_backends = wrapped_init_attention_backends
target.init_cuda_graphs = wrapped_init_cuda_graphs
class C_KVCacheConfiguratorHook(BaseHook):
HOOK_CLASS_NAME = "KVCacheConfigurator"
HOOK_MODULE_NAME = "sglang.srt.mem_cache.kv_cache_configurator"
@classmethod
def hook(cls, target):
original_configure = target.configure
original_init_pools = target._init_pools
supports_cpu_fp8_quant_method = hasattr(target, "_build_mha_quant_method")
def wrapped_configure(self, *args, **kwargs):
if not (
supports_cpu_fp8_quant_method
and getattr(self, "device", None) == "cpu"
and getattr(self, "kv_cache_dtype", None) == torch.float8_e4m3fn
):
return original_configure(self, *args, **kwargs)
# Newer SGLang runtimes validate CPU FP8 KV-cache support and select
# an AMX-only quant method. The simulator executes scheduler state on
# CPU while modeling the target accelerator's FP8 cache. Suppress the
# physical-CPU predicate only while native compact pools are built;
# all other platform capabilities and the logical dtype stay intact.
from sglang.srt.mem_cache.kv_cache_configurator import current_platform
original_is_cpu = current_platform.is_cpu
current_platform.is_cpu = lambda: False
try:
return original_configure(self, *args, **kwargs)
finally:
current_platform.is_cpu = original_is_cpu
def override_profile_available_bytes(self, pre_model_load_memory):
if self.server_args.max_total_tokens is not None:
from sglang.srt.model_executor.pool_configurator import (
create_memory_pool_configurator,
)
configurator = create_memory_pool_configurator(self)
target_tokens = self.server_args.max_total_tokens
page_size = _resolve_kv_page_size(self)
def resolved_tokens(budget_bytes):
try:
config = configurator.calculate_pool_sizes(
budget_bytes, page_size
)
except RuntimeError:
return 0
return config.max_total_num_tokens
lower, upper = 0, 1
while resolved_tokens(upper) < target_tokens:
lower, upper = upper, upper * 2
while lower + 1 < upper:
middle = (lower + upper) // 2
if resolved_tokens(middle) < target_tokens:
lower = middle
else:
upper = middle
return upper
model = ConfigManager.get_model_info()
if model is None:
model = resolve_model_info(self.model_config)
ConfigManager.set_model_info(model)
hardware = ConfigManager.get_accelerator_info()
scheduler_config = resolve_scheduler_config(
server_args=self.server_args,
model_config=self.model_config,
)
if hardware is None or scheduler_config is None:
raise RuntimeError(
"Simulator model, accelerator, and scheduler configuration "
"must be resolved before KV-cache pool sizing."
)
available_bytes = profile_device_available_bytes(
model=model,
device=hardware,
scheduler_config=scheduler_config,
)
if self.mambaish_config is not None:
rest_memory_gb = self._handle_max_mamba_cache(
available_bytes / (1 << 30)
)
available_bytes = int(rest_memory_gb * (1 << 30))
return available_bytes
def wrapped_init_pools(self, *args, **kwargs):
# Pool payload is never read during simulation. Preserve the native
# pool classes and allocator wiring, but allocate minimal payload
# dimensions and restore their logical metadata afterwards.
compact_attrs = (
"qk_nope_head_dim",
"qk_rope_head_dim",
"index_head_dim",
"kv_lora_rank",
"head_dim",
"v_head_dim",
"linear_value_head_dim",
"linear_key_head_dim",
"linear_conv_kernel_dim",
)
original_attrs = {
name: getattr(self.model_config, name)
for name in compact_attrs
if hasattr(self.model_config, name)
}
try:
for name in original_attrs:
setattr(self.model_config, name, 1)
pools = original_init_pools(self, *args, **kwargs)
finally:
for name, value in original_attrs.items():
setattr(self.model_config, name, value)
token_pool = pools.token_to_kv_pool
for name, value in original_attrs.items():
if hasattr(token_pool, name):
setattr(token_pool, name, value)
if (
hasattr(token_pool, "kv_cache_dim")
and token_pool.kv_cache_dim == 2
and "kv_lora_rank" in original_attrs
and "qk_rope_head_dim" in original_attrs
):
token_pool.kv_cache_dim = (
original_attrs["kv_lora_rank"] + original_attrs["qk_rope_head_dim"]
)
return pools
target.configure = wrapped_configure
target._profile_available_bytes = override_profile_available_bytes
target._init_pools = wrapped_init_pools
@@ -0,0 +1,22 @@
from sglang_simulator.simulation.types import RequestStats
class RequestStatsManager:
"""Shared request statistics manager for `Scheduler` and `HicacheController`."""
def __init__(self):
self.stats: dict[str, RequestStats] = {}
def get_req_stats(self, rid: str) -> RequestStats:
if rid not in self.stats:
self.stats[rid] = RequestStats(rid=rid)
return self.stats[rid]
def get_all_req_stats(self) -> list[RequestStats]:
return list(self.stats.values())
def reset(self):
self.stats.clear()
request_stats_manager = RequestStatsManager()
@@ -0,0 +1,648 @@
import heapq
import importlib
import json
import os
import time
from dataclasses import asdict
from typing import Any
from sglang_simulator.compat import validate_simulator_server_args
from sglang_simulator.hook import (
BaseHook,
is_class_hook_matched,
validate_required_class_hooks,
)
from sglang_simulator.hook.utils import get_obj_from_args
from sglang_simulator.simulation.manager import ConfigManager, Envs, StateManager
from sglang_simulator.simulation.sglang.req_stats_manager import request_stats_manager
from sglang_simulator.simulation.sglang.utils import (
resolve_model_info,
resolve_scheduler_config,
)
from sglang_simulator.simulation.types import (
RequestStats,
SimulationMode,
)
from sglang_simulator.simulation.utils import (
calc_iteration_metrics,
calc_metrics,
)
from sglang_simulator.time_predictor import InferTimePredictor
from sglang_simulator.time_predictor import ScheduleBatch as SimulationScheduleBatch
from sglang_simulator.time_predictor import ScheduleRequest
from sglang_simulator.utils import get_logger
from sglang_simulator.utils.json import CustomJsonEncoder
logger = get_logger("sgl_simulator")
def simulation_mode_log_message(mode: SimulationMode) -> str:
return f"SGLang Simulator simulation mode: {mode.value}"
def effective_l2_load_delay(
load_duration: float,
last_inference_duration: float,
overlap_schedule: bool,
) -> float:
if overlap_schedule:
return max(load_duration - last_inference_duration, 0.0)
return max(load_duration, 0.0)
def block_on_l2_load(mode: SimulationMode, delay: float) -> float:
"""Sleep for visible L2 load time and return actual blocked wall time."""
if mode != SimulationMode.BLOCKING or delay <= 0:
return 0.0
start = time.perf_counter()
time.sleep(delay)
return time.perf_counter() - start
class C_SglangPrefillAdderHook(BaseHook):
HOOK_CLASS_NAME = "PrefillAdder"
HOOK_MODULE_NAME = "sglang.srt.managers.schedule_policy"
@classmethod
def hook(cls, target):
original_add_one_req = target.add_one_req
def wrapped_add_one_req(self, *args, **kwargs):
req = get_obj_from_args(
"sglang.srt.managers.schedule_batch.Req",
*args,
**kwargs,
)
req_infos = request_stats_manager.get_req_stats(req.rid)
req_infos.before_adder_device_hit_len = len(req.prefix_indices)
req_infos.final_host_hit_len = req.host_hit_length
return original_add_one_req(self, *args, **kwargs)
target.add_one_req = wrapped_add_one_req
class ReqDispatcher:
_instance = None
_initialized = False
def __new__(cls, mode):
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self, mode: SimulationMode):
if self.__class__._initialized:
return
self.mode = mode
# If the simulation mode is `BLOCKING`, all requests are released immediately.
# If the simulation mode is `OFFLINE`, only control requests, such as `flush_cache`
# and `server_info`, are released immediately.
self.immediate_release_requests = []
self.future_queue: list[
tuple[float, int, Any]
] = [] # tuple(created time, salt, request)
self.offline_recv_all_requests = False
self.profile_active = False
@staticmethod
def simulation_created_time_s(simulation_args: dict) -> float:
if "created_time_ms" in simulation_args:
return simulation_args["created_time_ms"] / 1000.0
return simulation_args["created_time"]
def has_next(self) -> bool:
return len(self.future_queue) > 0
def next_req_from_future_ts(self) -> float:
return self.future_queue[0][0]
def reset(self) -> None:
self.immediate_release_requests.clear()
self.future_queue.clear()
self.offline_recv_all_requests = False
def add(self, reqs: list):
if self.mode == SimulationMode.BLOCKING:
self.immediate_release_requests.extend(reqs)
elif self.mode == SimulationMode.OFFLINE:
if self.offline_recv_all_requests:
self.immediate_release_requests.extend(reqs)
return
gen_requests = []
time.sleep(0.05) # waiting requests
for req in reqs:
if req.__class__.__name__ == "TokenizedGenerateReqInput":
gen_requests.append(req)
else:
# Such as: /profile_start, /flush_cache, etc.
self.immediate_release_requests.append(req)
# Add requests to future queue
for req in gen_requests:
sim_params = None
if req.sampling_params.custom_params is not None:
sim_params = req.sampling_params.custom_params.get("simulation")
if sim_params is None:
# There are some warm-up requests when starting the server without --skip-server-warmup.
self.immediate_release_requests.append(req)
logger.warning(
"Failed to extract the simulation parameters required for simulation from the request. Ignore this warning if the request is a warm-up request."
)
continue
if sim_params.get("queue_start"):
logger.debug(
"Add request to waiting queue with custom queue start timestamp."
)
self.future_queue.append(
(
sim_params.get("queue_start")
or self.simulation_created_time_s(sim_params),
time.time_ns(), # The request is not comparable, so add the salt to avoid comparison.
req,
)
)
if len(self.future_queue) != 0:
_, _, gen_req = self.future_queue[-1]
total_request = gen_req.sampling_params.custom_params["simulation"][
"total_request"
]
if len(self.future_queue) == total_request:
self.offline_recv_all_requests = True
heapq.heapify(self.future_queue)
logger.info("All requests received. Starting simulation now.")
else:
logger.info(
f"Offline simulation mode enabled. {total_request} requests expected in total. Received {len(self.future_queue)} requests so far."
)
def dispatch(self) -> list:
recv_reqs = []
recv_reqs.extend(self.immediate_release_requests)
self.immediate_release_requests.clear()
if self.mode == SimulationMode.OFFLINE and self.offline_recv_all_requests:
# Process the arrived requests only after all requests have been added to the future queue
current_timestamp = StateManager.get_global_clock()
while len(self.future_queue) > 0:
enqueue_time, _, req = self.future_queue[0]
if enqueue_time > current_timestamp:
break
recv_reqs.append(req)
heapq.heappop(self.future_queue)
now = time.time()
for req in recv_reqs:
if req.__class__.__name__ in [
"BatchTokenizedGenerateReqInput",
"TokenizedGenerateReqInput",
]:
simulation_args = None
if req.sampling_params.custom_params is not None:
simulation_args = req.sampling_params.custom_params.get(
"simulation"
)
# The warm-up request might not include any simulation arguments.
if simulation_args is None:
if self.mode != SimulationMode.BLOCKING or not self.profile_active:
continue
simulation_args = {}
req_stats = request_stats_manager.get_req_stats(req.rid)
req_stats.rid = req.rid
req_stats.input_length = len(req.input_ids)
req_stats.output_length = req.sampling_params.max_new_tokens
if self.mode == SimulationMode.BLOCKING:
req_stats.created_time = simulation_args.get(
"server_created_time", now
)
req_stats.last_event_time = req_stats.created_time
req_stats.queue_start = now
elif self.mode == SimulationMode.OFFLINE:
req_stats.created_time = self.simulation_created_time_s(
simulation_args
)
req_stats.last_event_time = req_stats.created_time
# Align with the real queue start timestamp if queue_start is not None. For debugging only.
queue_start = simulation_args.get("queue_start")
if queue_start is not None:
StateManager.set_global_clock(queue_start)
req_stats.queue_start = StateManager.get_global_clock()
if recv_reqs and StateManager.get_last_real_time_ts() == 0:
StateManager.set_last_real_time_ts(time.time())
StateManager.set_global_clock(
now if self.mode == SimulationMode.BLOCKING else 0
)
return recv_reqs
class C_SchedulerRequestReceiver(BaseHook):
HOOK_CLASS_NAME = "SchedulerRequestReceiver"
HOOK_MODULE_NAME = "sglang.srt.managers.scheduler_components.request_receiver"
# Older SGLang versions receive requests directly on Scheduler; that path is
# patched by C_SchedulerHook instead.
REQUIRED = False
REQ_DISPATCHER: ReqDispatcher = ReqDispatcher(
SimulationMode(Envs.simulation_mode())
)
@classmethod
def hook(cls, target):
original_recv_requests = target.recv_requests
def wrapped_recv_requests(self, *args, **kwargs):
recv_reqs = original_recv_requests(self, *args, **kwargs)
C_SchedulerRequestReceiver.REQ_DISPATCHER.add(recv_reqs)
return C_SchedulerRequestReceiver.REQ_DISPATCHER.dispatch()
target.recv_requests = wrapped_recv_requests
class C_SchedulerHook(BaseHook):
HOOK_CLASS_NAME = "Scheduler"
HOOK_MODULE_NAME = "sglang.srt.managers.scheduler"
INFERENCE_PREDICTOR: InferTimePredictor = None
ITERATION_STATS: list[dict] = []
TOTAL_PREDICTOR_TIME_COST = 0
GET_NEW_BATCH_PREFILL_TIME_COST = 0
SIMULATION_BATCH: SimulationScheduleBatch = None
OVERLAP_SCHEDULE: bool = False
SIM_MODE = SimulationMode(Envs.simulation_mode())
# Shared singleton instance with `C_SchedulerRequestReceiver.REQ_DISPATCHER`.
REQ_DISPATCHER = ReqDispatcher(SIM_MODE)
@classmethod
def hook(cls, target):
original_init = target.__init__
original_recv_requests = getattr(target, "recv_requests", None)
original_prefetch_kvcache = target._prefetch_kvcache
original_get_new_batch_prefill = target.get_new_batch_prefill
original_run_batch = target.run_batch
original_process_batch_result = target.process_batch_result
original_event_loop_normal = target.event_loop_normal
original_init_request_dispatcher = target.init_request_dispatcher
def override_event_loop_overlap(self, *args, **kwargs):
# To reduce the complexity of the simulation, the overlapping schedule is not needed.
return original_event_loop_normal(self, *args, **kwargs)
def wrapped_init(self, *args, **kwargs):
logger.info(simulation_mode_log_message(C_SchedulerHook.SIM_MODE))
# Supported entry points prepare the final config before publication.
server_args = get_obj_from_args(
"sglang.srt.server_args.ServerArgs", *args, **kwargs
)
validate_simulator_server_args(server_args)
C_SchedulerHook.OVERLAP_SCHEDULE = not getattr(
server_args, "disable_overlap_schedule", False
)
logger.debug(
f"Overlap schedule simulation mode: {C_SchedulerHook.OVERLAP_SCHEDULE}."
)
original_init(self, *args, **kwargs)
validate_required_class_hooks()
if original_recv_requests is None and not is_class_hook_matched(
C_SchedulerRequestReceiver
):
raise RuntimeError(
"SGLang Simulator could not hook a request receiver. The "
"simulator must be adapted to this SGLang revision."
)
try:
if ConfigManager.get_model_info() is None:
model = resolve_model_info(self.model_config)
ConfigManager.set_model_info(model)
model = ConfigManager.get_model_info()
hw = ConfigManager.get_accelerator_info()
if ConfigManager.get_scheduler_config() is None:
sched_config = resolve_scheduler_config(
server_args=self.server_args,
model_config=self.model_config,
)
ConfigManager.set_scheduler_config(sched_config)
sched_config = ConfigManager.get_scheduler_config()
C_SchedulerHook.INFERENCE_PREDICTOR = (
ConfigManager.get_inference_time_predictor(model, hw, sched_config)
)
except Exception as e:
logger.error(
f"Failed to initialize inference time predictor. Error: {e}"
)
raise e
def wrapped_recv_requests(self, *args, **kwargs) -> list:
recv_reqs = original_recv_requests(self, *args, **kwargs)
C_SchedulerHook.REQ_DISPATCHER.add(recv_reqs)
return C_SchedulerHook.REQ_DISPATCHER.dispatch()
def wrapped_get_new_batch_prefill(self, *args, **kwargs):
start = time.perf_counter()
result = original_get_new_batch_prefill(self, *args, **kwargs)
C_SchedulerHook.GET_NEW_BATCH_PREFILL_TIME_COST = (
time.perf_counter() - start
)
# Accept both a plan wrapper and a direct batch return value.
new_batch = getattr(result, "batch_to_run", result)
# A plan reports the running batch before self.running_batch is updated.
running_batch = getattr(result, "running_batch", self.running_batch)
now = time.time()
if new_batch is not None:
for req in new_batch.reqs:
req_stats = request_stats_manager.get_req_stats(req.rid)
req_stats.final_device_hit_len = req.cached_tokens
if req_stats.queue_end == -1:
if C_SchedulerHook.SIM_MODE == SimulationMode.BLOCKING:
req_stats.queue_end = now
else:
req_stats.queue_end = StateManager.get_global_clock()
else:
# Chunked request
pass
elif len(running_batch.reqs) == 0 and len(self.waiting_queue) > 0:
# Prefetching
StateManager.step_global_clock(0.005)
StateManager.set_current_inference_dur(0.005)
else:
# Idle stage, there are some requests pendding in the future queue.
if C_SchedulerHook.SIM_MODE == SimulationMode.OFFLINE and (
C_SchedulerHook.REQ_DISPATCHER.has_next()
and len(running_batch.reqs) == 0
):
next_created_time = (
C_SchedulerHook.REQ_DISPATCHER.next_req_from_future_ts()
)
StateManager.set_global_clock(next_created_time + 1e-6)
logger.debug(
f"Get new batch prefill: global iteration={StateManager.get_iteration()}, "
f"new batch={new_batch.batch_size() if new_batch is not None else 0}, "
f"waiting queue={len(self.waiting_queue)}"
)
return result
def wrapped_prefetch_kvcache(self, *args, **kwargs):
original_prefetch_kvcache(self, *args, **kwargs)
req = get_obj_from_args(
"sglang.srt.managers.schedule_batch.Req",
*args,
**kwargs,
)
req_stats = request_stats_manager.get_req_stats(req.rid)
req_stats.recv_device_hit_len = len(req.prefix_indices)
req_stats.recv_host_hit_len = req.host_hit_length
def wrapped_run_batch(self, *args, **kwargs):
ret = original_run_batch(self, *args, **kwargs)
batch = get_obj_from_args(
"sglang.srt.managers.schedule_batch.ScheduleBatch", *args, **kwargs
)
if ret.__class__.__name__ == "GenerationBatchResult":
simulation_batch = SimulationScheduleBatch(reqs=[])
if batch.forward_mode.is_extend():
for req in batch.reqs:
extend_length = getattr(req, "extend_input_len", None)
if extend_length is None:
# The range API represents extend tokens as a half-open interval.
extend_length = req.extend_range.length
simulation_batch.reqs.append(
ScheduleRequest(
extend_length=extend_length,
past_kv_length=len(req.prefix_indices)
+ len(req.output_ids),
)
)
elif batch.forward_mode.is_decode():
for req in batch.reqs:
simulation_batch.reqs.append(
ScheduleRequest(
extend_length=1,
past_kv_length=len(req.prefix_indices)
+ len(req.output_ids),
)
)
if not simulation_batch.is_empty():
StateManager.inc_iteration()
pred_start = time.perf_counter()
predicted_latency = (
C_SchedulerHook.INFERENCE_PREDICTOR.predict_infer_time(
simulation_batch
)
)
# Accumulate predictor execution time for performance analysis.
C_SchedulerHook.TOTAL_PREDICTOR_TIME_COST += (
time.perf_counter() - pred_start
)
predicted_latency = float(predicted_latency)
forward_latency = 0
if C_SchedulerHook.SIM_MODE == SimulationMode.BLOCKING:
time.sleep(abs(predicted_latency))
now = time.time()
forward_latency = now - StateManager.get_last_real_time_ts()
StateManager.set_last_real_time_ts(now)
else:
forward_latency = predicted_latency
StateManager.set_current_inference_dur(forward_latency)
C_SchedulerHook.SIMULATION_BATCH = simulation_batch
return ret
def wrapped_process_batch_result(self, *args, **kwargs):
process_batch_result_start = time.perf_counter()
ret = original_process_batch_result(self, *args, **kwargs)
process_batch_result_end = time.perf_counter()
batch = get_obj_from_args(
"sglang.srt.managers.schedule_batch.ScheduleBatch", *args, **kwargs
)
if batch is not None:
if len(batch.reqs) == 0:
return ret
hicache_l2_load_dur = StateManager.pop_hicache_l2_load_dur()
hicache_l2_load_stats = StateManager.pop_hicache_l2_load_stats()
hicache_l2_backup_dur = StateManager.pop_hicache_l2_backup_dur()
current_inference_dur = StateManager.get_current_inference_dur()
visible_l2_load_dur = effective_l2_load_delay(
hicache_l2_load_dur,
StateManager.get_last_inference_dur(),
C_SchedulerHook.OVERLAP_SCHEDULE,
)
blocked_l2_wall_dur = block_on_l2_load(
C_SchedulerHook.SIM_MODE,
visible_l2_load_dur,
)
StateManager.step_global_clock(visible_l2_load_dur)
StateManager.step_global_clock(current_inference_dur)
# Step CPU overhead BEFORE recording latencies,
# so current iter's CPU time is reflected in current iter's TTFT.
now = time.time()
cpu_overhead = max(
now - StateManager.get_last_real_time_ts() - blocked_l2_wall_dur,
0.0,
)
StateManager.step_global_clock(cpu_overhead)
StateManager.set_last_real_time_ts(now)
request_response_time = StateManager.get_global_clock()
# Request statistics
for req in batch.reqs:
if len(req.output_ids) != 0: # not chunked
req_stats = request_stats_manager.get_req_stats(req.rid)
req_stats.gen_token_latencies.append(
request_response_time
- req_stats.last_event_time # queue duration
)
req_stats.last_event_time = request_response_time
else:
# Chunked request: nothing to do
pass
# Iteration statistics
C_SchedulerHook.ITERATION_STATS.append(
{
"requests": C_SchedulerHook.SIMULATION_BATCH.request_info(),
"forward_latency": current_inference_dur,
"l2_load_latency": hicache_l2_load_dur,
"l2_blocking_wall_latency": blocked_l2_wall_dur,
**hicache_l2_load_stats,
"l2_backup_latency": hicache_l2_backup_dur,
"preprocess_latency": C_SchedulerHook.GET_NEW_BATCH_PREFILL_TIME_COST,
"postprocess_latency": process_batch_result_end
- process_batch_result_start,
"cpu_overhead": cpu_overhead,
}
)
else:
now = time.time()
StateManager.step_global_clock(
now - StateManager.get_last_real_time_ts()
)
StateManager.set_last_real_time_ts(now)
return ret
def override_profile(req, *args, **kwargs):
is_start_profile = req.req_type.name == "START_PROFILE"
stats: list[RequestStats] = []
for item in request_stats_manager.get_all_req_stats():
if item.rid is not None and item.input_length > 0:
stats.append(item)
stats = sorted(stats, key=lambda req: req.created_time)
output_dir = Envs.output_dir()
os.makedirs(output_dir, exist_ok=True)
if len(stats) > 0:
min_created_time = stats[0].created_time
# Align timestamps
for item in stats:
item.created_time -= min_created_time
item.queue_start -= min_created_time
item.queue_end -= min_created_time
item.last_event_time -= min_created_time
metrics = calc_metrics(stats)
metrics["time_cost"] = (
time.time() - StateManager.get_last_flush_time_ts()
)
metrics["predictor_time_cost"] = (
C_SchedulerHook.TOTAL_PREDICTOR_TIME_COST
)
metrics.update(
calc_iteration_metrics(C_SchedulerHook.ITERATION_STATS, metrics)
)
metrics.update(C_SchedulerHook.INFERENCE_PREDICTOR.get_metrics())
try:
with open(f"{output_dir}/metrics.json", "w") as f:
f.write(json.dumps(metrics, cls=CustomJsonEncoder) + "\n")
with open(f"{output_dir}/iteration.jsonl", "w") as f:
for item in C_SchedulerHook.ITERATION_STATS:
f.write(json.dumps(item) + "\n")
with open(f"{output_dir}/request.jsonl", "w") as f:
for item in stats:
f.write(json.dumps(asdict(item)) + "\n")
logger.info(f"Simulation results saved to {output_dir}.")
except Exception as e:
logger.error(f"Failed to dump results. Error: {e}")
else:
logger.warning("No request statistics available.")
StateManager.reset()
StateManager.set_last_flush_time_ts(time.time())
request_stats_manager.reset()
C_SchedulerHook.ITERATION_STATS.clear()
C_SchedulerHook.TOTAL_PREDICTOR_TIME_COST = 0
C_SchedulerHook.REQ_DISPATCHER.reset()
C_SchedulerHook.REQ_DISPATCHER.profile_active = is_start_profile
C_SchedulerHook.INFERENCE_PREDICTOR.reset_metrics()
ProfileReqOutput = getattr(
importlib.import_module("sglang.srt.managers.io_struct"),
"ProfileReqOutput",
)
result = {
"total_request": len(stats),
"output_directory": output_dir,
}
return ProfileReqOutput(
success=True,
message=json.dumps(result),
)
def wrapped_init_request_dispatcher(self, *args, **kwargs):
ret = original_init_request_dispatcher(self, *args, **kwargs)
_request_dispatcher = getattr(self, "_request_dispatcher", None)
if _request_dispatcher is not None:
for ty in _request_dispatcher._mapping.keys():
if ty.__name__ == "ProfileReq":
_request_dispatcher._mapping[ty] = override_profile
return ret
target.event_loop_overlap = override_event_loop_overlap
target.__init__ = wrapped_init
target.get_new_batch_prefill = wrapped_get_new_batch_prefill
target.run_batch = wrapped_run_batch
target.process_batch_result = wrapped_process_batch_result
target._prefetch_kvcache = wrapped_prefetch_kvcache
target.init_request_dispatcher = wrapped_init_request_dispatcher
if original_recv_requests:
target.recv_requests = wrapped_recv_requests
@@ -0,0 +1,15 @@
import sys
import types
def install_load_utils_stub() -> None:
"""Install the kernel loader stub before importing the sgl_kernel package."""
module_name = "sgl_kernel.load_utils"
module = sys.modules.get(module_name)
if module is None:
module = types.ModuleType(module_name)
module.__package__ = "sgl_kernel"
sys.modules[module_name] = module
module._load_architecture_specific_ops = lambda *args, **kwargs: None
module._preload_cuda_library = lambda *args, **kwargs: None
@@ -0,0 +1,34 @@
from sglang_simulator.hook import BaseHook
class C_UnifiedRadixCacheHook(BaseHook):
"""Drive Unified HiCache storage work from the simulator's logical clock."""
HOOK_CLASS_NAME = "UnifiedRadixCache"
HOOK_MODULE_NAME = "sglang.srt.mem_cache.unified_radix_cache"
REQUIRED = False
@classmethod
def hook(cls, target):
original_check_hicache_events = target.check_hicache_events
def handle_pending_operations(controller):
if controller is None:
return
backup_handler = getattr(controller, "handle_backup_operation", None)
prefetch_handler = getattr(controller, "handle_prefetch_operation", None)
if backup_handler is not None:
backup_handler()
if prefetch_handler is not None:
prefetch_handler()
def wrapped_check_hicache_events(self, *args, **kwargs):
controller = getattr(self, "cache_controller", None)
handle_pending_operations(controller)
result = original_check_hicache_events(self, *args, **kwargs)
# Unified allocates host pages while draining its scheduler-side
# control queues. Process those newly admitted reads immediately.
handle_pending_operations(controller)
return result
target.check_hicache_events = wrapped_check_hicache_events
@@ -0,0 +1,130 @@
import typing
from sglang_simulator.simulation.types import SchedulerConfig
from sglang_simulator.spec import DataType, ModelInfo
if typing.TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
def _resolve_model_config(server_args: "ServerArgs", model_config=None):
if model_config is not None:
return model_config
get_model_config = getattr(server_args, "get_model_config", None)
if callable(get_model_config):
return get_model_config()
return server_args.model_config
def _resolved_server_args(server_args: "ServerArgs") -> dict:
"""Return effective ServerArgs values when the runtime exposes them."""
resolved_dict = getattr(server_args, "resolved_dict", None)
if not callable(resolved_dict):
return {}
try:
values = resolved_dict()
except (AttributeError, RuntimeError, TypeError):
return {}
return values if isinstance(values, dict) else {}
def resolve_scheduler_config(
server_args: "ServerArgs",
model_config: typing.Optional["ModelConfig"] = None,
) -> SchedulerConfig:
from sglang.version import __version__
resolved = _resolved_server_args(server_args)
def get_arg(name: str, default=None):
value = resolved.get(name)
if value is None:
value = getattr(server_args, name, None)
return default if value is None else value
dtype = get_arg("dtype", "auto")
if dtype == "auto":
model_config = _resolve_model_config(server_args, model_config)
dtype = str(model_config.dtype).strip("torch.")
data_type = DataType.from_torch_dtype(dtype)
return SchedulerConfig(
data_type=data_type,
kv_cache_data_type=DataType.from_torch_dtype(get_arg("kv_cache_dtype"))
or data_type,
mem_fraction_static=get_arg("mem_fraction_static"),
max_total_tokens=get_arg("max_total_tokens"),
tp_size=get_arg("tp_size"),
ep_size=get_arg("ep_size"),
dp_size=get_arg("dp_size"),
pp_size=get_arg("pp_size"),
cp_size=get_arg("attn_cp_size", 1),
cp_style=get_arg("cp_style", "none"),
page_size=get_arg("page_size"),
swa_full_tokens_ratio=get_arg("swa_full_tokens_ratio"),
kv_bytes_per_token_per_gpu=get_arg("kv_bytes_per_token_per_gpu"),
hicache_ratio=get_arg("hicache_ratio"),
enable_hierarchical_cache=get_arg("enable_hierarchical_cache"),
backend_name="sglang",
backend_version=__version__,
)
def resolve_model_info(model_config: "ModelConfig") -> ModelInfo:
from sglang.srt.configs.model_config import AttentionArch
torch_dtype = str(model_config.dtype).strip("torch.")
if model_config.attention_arch == AttentionArch.MHA:
return ModelInfo(
hf_config=model_config.hf_text_config,
model_path=model_config.model_path,
attention_arch="MHA",
context_len=model_config.context_len,
hidden_size=model_config.hidden_size,
head_dim=model_config.head_dim,
num_attention_heads=model_config.num_attention_heads,
num_hidden_layers=model_config.num_hidden_layers,
num_key_value_heads=model_config.num_key_value_heads,
v_head_dim=model_config.v_head_dim,
vocab_size=model_config.vocab_size,
# DSv4-style models (e.g. DSv4-Pro) report attention_arch=MHA because
# sglang routes them through a custom `attention_backend='dsv4'`, not
# MLA. But they still carry compress_ratios + indexer + SWA fields on
# ModelConfig, and is_dsv4() needs them to take the right calculator
# branch. getattr makes this a no-op for true MHA models.
compression_ratios=getattr(model_config, "compress_ratios", None),
indexer_head_dim=getattr(model_config, "index_head_dim", None),
window_size=getattr(model_config, "window_size", None),
qk_nope_head_dim=getattr(model_config, "qk_nope_head_dim", None),
qk_rope_head_dim=getattr(model_config, "qk_rope_head_dim", None),
torch_dtype=torch_dtype,
)
elif model_config.attention_arch == AttentionArch.MLA:
return ModelInfo(
hf_config=model_config.hf_text_config,
model_path=model_config.model_path,
attention_arch="MLA",
context_len=model_config.context_len,
hidden_size=model_config.hidden_size,
head_dim=model_config.head_dim,
num_attention_heads=model_config.num_attention_heads,
num_hidden_layers=model_config.num_hidden_layers,
num_key_value_heads=model_config.num_key_value_heads,
v_head_dim=model_config.v_head_dim,
vocab_size=model_config.vocab_size,
qk_rope_head_dim=model_config.qk_rope_head_dim,
qk_nope_head_dim=model_config.qk_nope_head_dim,
kv_lora_rank=model_config.kv_lora_rank,
compression_ratios=getattr(model_config, "compress_ratios", None),
indexer_head_dim=getattr(model_config, "index_head_dim", None),
window_size=getattr(model_config, "window_size", None),
torch_dtype=torch_dtype,
)
else:
raise ValueError(
f"The attention type of `{model_config.attention_arch}` is not supported now."
)
@@ -0,0 +1,134 @@
from dataclasses import dataclass, field
from enum import Enum
from typing import Optional, Union
from sglang_simulator.spec import AcceleratorInfo, DataType
@dataclass
class SchedulerConfig:
data_type: Optional[DataType] = (
None # Data type for model weights and activations. If none is set, it will be automatically detected.
)
kv_cache_data_type: Optional[DataType] = None
# AIC adapter overrides — bypass MAP_DTYPE_TO_* lookup when set.
# Pass aiconfigurator MoEQuantMode/FMHAQuantMode/CommQuantMode enum name as string
# (e.g. 'w4a8_mxfp4_mxfp8' for DSv4-Pro on Blackwell).
moe_quant_mode_override: Optional[str] = None
fmha_quant_mode_override: Optional[str] = None
comm_quant_mode_override: Optional[str] = None
mem_fraction_static: Optional[float] = None
max_total_tokens: Optional[int] = None
tp_size: int = 1
ep_size: int = 1
dp_size: int = 1
pp_size: int = 1
cp_size: int = 1
cp_style: str = "none"
# DSv4 KV cache calculator inputs (sourced from server_args)
page_size: Optional[int] = None
swa_full_tokens_ratio: Optional[float] = None
# Optional explicit override of per-GPU KV bytes/token, sourced from
# sglang server startup log: "KV Cache is allocated. #tokens: N, KV size: G GB"
# kv_bytes_per_token_per_gpu = G * 1024**3 / N
# When set, takes priority over the model-derived calculator path.
# Useful for models where sglang doesn't expose its KV calculator output
# (e.g. GlmMoeDsa) and we want metrics to match the live sglang server.
kv_bytes_per_token_per_gpu: Optional[float] = None
# L2 host KV pool sizing: host_pool_tokens = hicache_ratio * max_total_tokens.
hicache_ratio: Optional[float] = None
enable_hierarchical_cache: Optional[bool] = None
# framework backend
backend_name: str = "sglang"
backend_version: Optional[str] = None
@property
def attn_tp_size(self) -> int:
divisor = self.dp_size * self.cp_size
if self.tp_size % divisor != 0:
raise ValueError(
"tp_size must be divisible by dp_size * cp_size: "
f"{self.tp_size} % ({self.dp_size} * {self.cp_size}) != 0"
)
return self.tp_size // divisor
@property
def attn_dp_size(self) -> int:
return self.dp_size
@property
def moe_tp_size(self) -> int:
if self.tp_size % self.ep_size != 0:
raise ValueError(
"tp_size must be divisible by ep_size: "
f"{self.tp_size} % {self.ep_size} != 0"
)
return self.tp_size // self.ep_size
@property
def moe_ep_size(self) -> int:
return self.ep_size
class SimulationMode(Enum):
BLOCKING = "BLOCKING"
OFFLINE = "OFFLINE"
@dataclass(slots=True)
class RequestStats:
rid: str = ""
last_event_time: float = 0.0
input_length: int = 1
output_length: int = 1
# Prefix cache stats
recv_device_hit_len: int = 0
# Device hit length before `get_new_batch_prefill`.
# It may decrease if queued requests trigger KV eviction.
before_adder_device_hit_len: int = 0
final_device_hit_len: int = 0
recv_host_hit_len: int = 0 # Host hit length before prefetch
final_host_hit_len: int = 0 # Host hit length after prefetch
recv_storage_hit_len: int = 0 # Storage hit length at prefetch enqueue
final_storage_hit_len: int = 0 # Storage hit length at prefetch end
queue_start: float = -1
queue_end: float = -1
created_time: float = -1
gen_token_latencies: list[float] = field(default_factory=list)
def is_complete(self) -> bool:
return True
def _bandwidth_property(gb_attr: str):
def getter(self):
gb_value = getattr(self, gb_attr)
return gb_value * 1e9 if gb_value else None
return property(getter)
@dataclass
class PlatformConfig:
device: Union[AcceleratorInfo, str]
# Storage configuration for hierarchical cache management.
disk_capacity_gb: Optional[float] = None
disk_read_bandwidth_gb: Optional[float] = None
disk_write_bandwidth_gb: Optional[float] = None
memory_capacity_gb: Optional[float] = None
memory_read_bandwidth_gb: Optional[float] = None
memory_write_bandwidth_gb: Optional[float] = None
num_device_per_node: int = 8
# Bandwidth properties (in bytes, converted from GB)
disk_read_bandwidth = _bandwidth_property("disk_read_bandwidth_gb")
disk_write_bandwidth = _bandwidth_property("disk_write_bandwidth_gb")
memory_read_bandwidth = _bandwidth_property("memory_read_bandwidth_gb")
memory_write_bandwidth = _bandwidth_property("memory_write_bandwidth_gb")
@@ -0,0 +1,239 @@
import numpy as np
from sglang_simulator.simulation.types import RequestStats, SchedulerConfig
from sglang_simulator.spec.accelerator import AcceleratorInfo
from sglang_simulator.spec.model import ModelInfo
from sglang_simulator.time_predictor.aiconfigurator import get_perf_model
def calc_kv_cache_cell_elems(model_info: ModelInfo, tp_size: int, pp_size: int) -> int:
num_layers = model_info.num_hidden_layers // pp_size
if model_info.is_mla():
return (model_info.kv_lora_rank + model_info.qk_rope_head_dim) * num_layers
else:
num_kv_heads = max(model_info.num_key_value_heads // tp_size, 1)
return num_kv_heads * model_info.head_dim * num_layers * 2
def calc_kv_cache_per_layer_elems(
model_info: ModelInfo, tp_size: int, pp_size: int
) -> int:
if model_info.is_mla():
return model_info.kv_lora_rank + model_info.qk_rope_head_dim
else:
num_kv_heads = max(model_info.num_key_value_heads // tp_size, 1)
return num_kv_heads * model_info.head_dim * 2
def profile_device_available_bytes(
model: ModelInfo, device: AcceleratorInfo, scheduler_config: SchedulerConfig
) -> int:
"""Return the simulated per-GPU byte budget available to KV-cache pools."""
# Simulation capacity must come from the declared target accelerator. Do
# not fall back to the local CUDA device: doing so would make an identical
# simulation config host-dependent and could silently simulate the wrong
# hardware.
if device.hbm_capacity_gb is None:
raise ValueError(
"Cannot estimate max_total_num_tokens: the simulated accelerator "
f"{device.name!r} has no hbm_capacity_gb. Add the accelerator to "
"the simulator hardware registry, provide hbm_capacity_gb in the "
"simulation config, or set max_total_tokens explicitly. The "
"simulator never falls back to the local GPU memory capacity."
)
perf_model = get_perf_model(scheduler_config, model)
weights = 0
for op in perf_model.context_ops:
weights += op.get_weights()
# Count weights on a single GPU
weights /= perf_model.config.pp_size
framework_reserved_mem_gb = 1.4
rest_memory = (
scheduler_config.mem_fraction_static * device.hbm_capacity_gb
- framework_reserved_mem_gb
) * (1 << 30) - weights
return int(rest_memory)
def calc_input_token_metrics(
total_input: int,
total_reused_tokens: int,
total_dur_s: float,
) -> dict:
"""Compute model-independent new-input token count and throughput."""
dur_s = max(total_dur_s, 1e-9)
total_new_input_tokens = total_input - total_reused_tokens
new_input_write_thr_tokens = total_new_input_tokens / dur_s
return {
"total_new_input": total_new_input_tokens,
"new_input_write_throughput_tokens_per_s": new_input_write_thr_tokens,
}
def calc_iteration_metrics(
iteration_stats: list[dict], request_metrics: dict | None = None
) -> dict:
"""Aggregate per-iteration simulator latency into result metrics."""
iterations = len(iteration_stats)
forward_s = sum(
float(item.get("forward_latency", 0) or 0) for item in iteration_stats
)
l2_load_s = sum(
float(item.get("l2_load_latency", 0) or 0) for item in iteration_stats
)
cpu_s = sum(float(item.get("cpu_overhead", 0) or 0) for item in iteration_stats)
total_s = forward_s + l2_load_s + cpu_s
avg_iter_latency_ms = total_s / iterations * 1000 if iterations else 0
metrics = {
"iterations": iterations,
"avg_iter_latency_ms": avg_iter_latency_ms,
}
if request_metrics:
mean_ttft_ms = request_metrics.get("mean_ttft_ms")
mean_queue_ms = request_metrics.get("mean_queue_ms")
if mean_ttft_ms is not None and mean_queue_ms is not None:
mean_exec_ms = mean_ttft_ms - mean_queue_ms
metrics["mean_exec_ms"] = mean_exec_ms
metrics["avg_iters_per_req"] = (
mean_exec_ms / avg_iter_latency_ms if avg_iter_latency_ms else None
)
return metrics
def calc_metrics(requests: list[RequestStats]) -> dict:
ttfts = []
tpots = []
itls = []
e2e_latencies = []
total_dur_s = 1e-9
total_input = 0
total_output = 0
completed = 0
total_reused_tokens = 0
total_device_hit_tokens = 0
total_host_hit_tokens = 0
total_storage_hit_tokens = 0
queue_durs = []
dispatch_wait_durs = []
arrival_to_prefill_durs = []
output_token_timestamps = []
concurrency_events = []
for req in requests:
if not req.is_complete():
continue
completed += 1
ttfts.append(req.gen_token_latencies[0])
# Queue latency is the time spent in SGLang's waiting queue before
# the request's first prefill admission.
queue_durs.append(req.queue_end - req.queue_start)
dispatch_wait_durs.append(req.queue_start - req.created_time)
arrival_to_prefill_durs.append(req.queue_end - req.created_time)
if len(req.gen_token_latencies) > 1:
# output length > 1
tpots.append(np.mean(req.gen_token_latencies[1:]))
itls.extend(req.gen_token_latencies[1:])
e2e_latencies.append(sum(req.gen_token_latencies))
token_timestamp = req.created_time
for token_latency in req.gen_token_latencies:
token_timestamp += token_latency
output_token_timestamps.append(token_timestamp)
concurrency_events.append((req.created_time, 1))
concurrency_events.append((req.last_event_time, -1))
total_dur_s = max(total_dur_s, req.last_event_time)
total_input += req.input_length
total_output += req.output_length
total_reused_tokens += req.final_device_hit_len
total_device_hit_tokens += req.final_device_hit_len - req.final_host_hit_len
total_host_hit_tokens += req.final_host_hit_len - req.final_storage_hit_len
total_storage_hit_tokens += req.final_storage_hit_len
input_token_metrics = calc_input_token_metrics(
total_input=total_input,
total_reused_tokens=total_reused_tokens,
total_dur_s=total_dur_s,
)
max_output_tokens_per_s = 0.0
if output_token_timestamps:
first_created_time = min(
req.created_time for req in requests if req.is_complete()
)
num_buckets = int(max(output_token_timestamps) - first_created_time) + 1
output_tokens_per_s = np.zeros(max(num_buckets, 1))
for timestamp in output_token_timestamps:
bucket = int(timestamp - first_created_time)
output_tokens_per_s[bucket] += 1
max_output_tokens_per_s = float(np.max(output_tokens_per_s))
max_concurrent_requests = 0
current_concurrent_requests = 0
# Treat request intervals as [created_time, last_event_time): requests that
# finish exactly when another arrives are not simultaneously active.
for _, delta in sorted(concurrency_events, key=lambda event: (event[0], event[1])):
current_concurrent_requests += delta
max_concurrent_requests = max(
max_concurrent_requests, current_concurrent_requests
)
return {
"num_requests": len(requests),
"completed": completed,
"total_input": total_input,
"total_output": total_output,
"duration": total_dur_s,
"request_throughput": completed / total_dur_s,
"input_throughput": total_input / total_dur_s,
"output_throughput": total_output / total_dur_s,
"total_throughput": (total_input + total_output) / total_dur_s,
"prefix_cache_reused_ratio": (
0 if total_input == 0 else total_reused_tokens / total_input
),
"kv_cache_storage_hit_ratio": (
0 if total_input == 0 else total_storage_hit_tokens / total_input
),
"kv_cache_host_hit_ratio": (
0 if total_input == 0 else total_host_hit_tokens / total_input
),
"kv_cache_device_hit_ratio": (
0 if total_input == 0 else total_device_hit_tokens / total_input
),
**input_token_metrics,
"mean_ttft_ms": np.mean(ttfts or 0) * 1000,
"median_ttft_ms": np.median(ttfts or 0) * 1000,
"std_ttft_ms": np.std(ttfts or 0) * 1000,
"p90_ttft_ms": np.percentile(ttfts or 0, 90) * 1000,
"p95_ttft_ms": np.percentile(ttfts or 0, 95) * 1000,
"p99_ttft_ms": np.percentile(ttfts or 0, 99) * 1000,
"mean_queue_ms": max(np.mean(queue_durs or 0), 0.0) * 1000,
"mean_dispatch_wait_ms": (max(np.mean(dispatch_wait_durs or 0), 0.0) * 1000),
"mean_arrival_to_prefill_ms": (
max(np.mean(arrival_to_prefill_durs or 0), 0.0) * 1000
),
"mean_tpot_ms": np.mean(tpots or 0) * 1000,
"median_tpot_ms": np.median(tpots or 0) * 1000,
"std_tpot_ms": np.std(tpots or 0) * 1000,
"p90_tpot_ms": np.percentile(tpots or 0, 90) * 1000,
"p95_tpot_ms": np.percentile(tpots or 0, 95) * 1000,
"p99_tpot_ms": np.percentile(tpots or 0, 99) * 1000,
"mean_itl_ms": np.mean(itls or 0) * 1000,
"median_itl_ms": np.median(itls or 0) * 1000,
"std_itl_ms": np.std(itls or 0) * 1000,
"p90_itl_ms": np.percentile(itls or 0, 90) * 1000,
"p95_itl_ms": np.percentile(itls or 0, 95) * 1000,
"p99_itl_ms": np.percentile(itls or 0, 99) * 1000,
"max_itl_ms": np.max(itls or 0) * 1000,
"mean_e2e_latency_ms": np.mean(e2e_latencies or 0) * 1000,
"median_e2e_latency_ms": np.median(e2e_latencies or 0) * 1000,
"std_e2e_latency_ms": np.std(e2e_latencies or 0) * 1000,
"p90_e2e_latency_ms": np.percentile(e2e_latencies or 0, 90) * 1000,
"p95_e2e_latency_ms": np.percentile(e2e_latencies or 0, 95) * 1000,
"p99_e2e_latency_ms": np.percentile(e2e_latencies or 0, 99) * 1000,
"concurrency": np.sum(e2e_latencies or 0) / total_dur_s,
"max_output_tokens_per_s": max_output_tokens_per_s,
"max_concurrent_requests": max_concurrent_requests,
"time_cost": -1, # Updated by external benchmark caller
}
@@ -0,0 +1,5 @@
from sglang_simulator.spec.accelerator import AcceleratorInfo
from sglang_simulator.spec.data_type import DataType
from sglang_simulator.spec.model import ModelInfo
__all__ = ["AcceleratorInfo", "ModelInfo", "DataType"]
@@ -0,0 +1,4 @@
from sglang_simulator.spec.accelerator.base import AcceleratorInfo
from sglang_simulator.spec.accelerator.info import NVIDIA
__all__ = ["AcceleratorInfo", "NVIDIA"]
@@ -0,0 +1,100 @@
from copy import deepcopy
from dataclasses import dataclass, field
from typing import Dict, Optional, Union
from sglang_simulator.spec.data_type import DataType
from sglang_simulator.utils import get_logger
_all_accs_: Dict[str, "AcceleratorInfo"] = {}
_acc_alias: Dict[str, str] = {}
logger = get_logger("sgl_simulator")
@dataclass
class AcceleratorInfo:
name: str
vendor: str
hbm_capacity_gb: int
hbm_bandwidth_gb: int
intra_node_bandwidth_gb: Optional[int] = None # scale up
inter_node_bandwidth_gb: int = 64 # scale out
device_alias: list = field(default_factory=list)
tflops: dict = field(default_factory=dict)
ref: str = ""
@classmethod
def from_dict(cls, config: Dict, save_to_registry: bool = False):
acc = cls(**config)
if save_to_registry:
if acc.name in _acc_alias:
logger.error(f"{acc.name} is already in registry")
_all_accs_[acc.name.upper()] = acc
for alias in acc.device_alias:
if alias in _acc_alias:
logger.warning(f"Device alias [{alias}] is already in registry.")
else:
_acc_alias[alias] = acc.name.upper()
return acc
def flops(self, datatype: Union[str, DataType] = DataType.FP16):
if isinstance(datatype, DataType):
datatype = datatype.value
return self.tflops.get(datatype, 1) * 1e12
def tensor_flops(self, datatype: Union[str, DataType] = DataType.FP16_TENSOR):
if isinstance(datatype, DataType):
datatype = datatype.value
if not datatype.endswith(DataType.tensor_suffix()):
datatype += DataType.tensor_suffix()
tflops = self.tflops.get(datatype, None)
return None if tflops is None else tflops * 1e12
@property
def hbm_io_bw(self):
return self.hbm_bandwidth_gb * 1e9
@property
def hbm_bytes(self):
return self.hbm_capacity_gb * 1e9
@property
def intra_node_bw(self) -> Optional[float]:
if self.intra_node_bandwidth_gb is None:
return None
return self.intra_node_bandwidth_gb * 1e9
@property
def inter_node_bw(self):
return self.inter_node_bandwidth_gb * 1e9
@staticmethod
def find_by_hw_name(hw_name: str) -> Union[None, "AcceleratorInfo"]:
if hw_name in _acc_alias:
hw = _all_accs_.get(_acc_alias[hw_name], None)
if hw is not None:
hw = deepcopy(hw)
hw.name = hw_name
return hw
else:
return _all_accs_.get(hw_name.upper(), None)
@staticmethod
def list_all_hws() -> Dict[str, "AcceleratorInfo"]:
return _all_accs_
@classmethod
def from_config(cls, config: Dict):
hw_info = cls.find_by_hw_name(config["name"])
return cls(**config) if hw_info is None else hw_info
def __eq__(self, value):
if isinstance(value, str):
value = self.find_by_hw_name(value)
if isinstance(value, AcceleratorInfo):
return _acc_alias.get(value.name, value.name.upper()) == _acc_alias.get(
self.name, self.name.upper()
)
return False
@@ -0,0 +1,24 @@
from sglang_simulator.spec.accelerator.base import AcceleratorInfo
class NVIDIA:
NVIDIA_H20 = AcceleratorInfo.from_dict(
config={
"name": "NVIDIA H20",
"device_alias": ["H20", "h20_sxm"],
"tflops": {
"FP8_TENSOR": 296,
"INT8_TENSOR": 296,
"FP16_TENSOR": 148,
"BF16_TENSOR": 148,
"FP32": 74,
},
"hbm_capacity_gb": 96,
"hbm_bandwidth_gb": 4022,
"inter_node_bandwidth_gb": 64,
"intra_node_bandwidth_gb": 450,
"vendor": "NVIDIA",
"ref": "https://viperatech.com/product/nvidia-hgx-h20",
},
save_to_registry=True,
)
@@ -0,0 +1,104 @@
from enum import Enum, unique
from typing import Dict, Optional
_BYTES_MAP: dict["DataType", float] = {}
_ALIAS_MAP: Dict[str, str] = {}
_TORCH_DTYPE_TO_DATA_TYPE: Dict[str, "DataType"] = {}
@unique
class DataType(Enum):
INT4 = "INT4"
INT8 = "INT8"
INT16 = "INT16"
INT32 = "INT32"
INT64 = "INT64"
FP4 = "FP4"
FP8 = "FP8"
FP16 = "FP16"
BF16 = "BF16"
TF32 = "TF32"
FP32 = "FP32"
FP64 = "FP64"
# tensor
INT4_TENSOR = "INT4_TENSOR"
INT8_TENSOR = "INT8_TENSOR"
INT16_TENSOR = "INT16_TENSOR"
INT32_TENSOR = "INT32_TENSOR"
INT64_TENSOR = "INT64_TENSOR"
FP4_TENSOR = "FP4_TENSOR"
FP8_TENSOR = "FP8_TENSOR"
FP16_TENSOR = "FP16_TENSOR"
BF16_TENSOR = "BF16_TENSOR"
TF32_TENSOR = "TF32_TENSOR"
FP32_TENSOR = "FP32_TENSOR"
FP64_TENSOR = "FP64_TENSOR"
# FIXME: This map will be added as a enum member.
@property
def bytes(self) -> float:
return _BYTES_MAP.get(self, 1)
@classmethod
def tensor_suffix(cls) -> str:
return "_TENSOR"
@classmethod
def alias(cls):
return _ALIAS_MAP
@classmethod
def from_torch_dtype(cls, dtype: str) -> Optional["DataType"]:
return _TORCH_DTYPE_TO_DATA_TYPE.get(dtype.lower())
_BYTES_MAP.update(
{
DataType.INT4: 0.5,
DataType.INT8: 1,
DataType.INT16: 2,
DataType.INT32: 4,
DataType.INT64: 8,
DataType.FP4: 0.5,
DataType.FP8: 1,
DataType.FP16: 2,
DataType.BF16: 2,
DataType.TF32: 4,
DataType.FP32: 4,
DataType.FP64: 8,
DataType.INT4_TENSOR: 0.5,
DataType.INT8_TENSOR: 1,
DataType.INT16_TENSOR: 2,
DataType.INT32_TENSOR: 4,
DataType.INT64_TENSOR: 8,
DataType.FP4_TENSOR: 0.5,
DataType.FP8_TENSOR: 1,
DataType.FP16_TENSOR: 2,
DataType.BF16_TENSOR: 2,
DataType.TF32_TENSOR: 4,
DataType.FP32_TENSOR: 4,
DataType.FP64_TENSOR: 8,
}
)
_ALIAS_MAP.update(
{
"int8": "INT8",
"float8": "FP8",
"float16": "FP16",
"float32": "FP32",
"bfloat16": "BF16",
}
)
_TORCH_DTYPE_TO_DATA_TYPE.update(
{
"fp8": DataType.FP8,
"int8": DataType.INT8,
"float8": DataType.FP8,
"float16": DataType.FP16,
"float32": DataType.FP32,
"bfloat16": DataType.BF16,
}
)
@@ -0,0 +1,3 @@
from sglang_simulator.spec.model.base import ModelInfo
__all__ = ["ModelInfo"]
@@ -0,0 +1,44 @@
from dataclasses import dataclass
from typing import Optional
from sglang_simulator.utils import get_logger
logger = get_logger("sgl_simulator")
@dataclass
class ModelInfo:
hf_config: Optional[dict] = None
model_path: Optional[str] = None
attention_arch: Optional[str] = None # MLA | MHA
context_len: Optional[int] = None
hidden_size: Optional[int] = None
head_dim: Optional[int] = None
num_attention_heads: Optional[int] = None
num_hidden_layers: Optional[int] = None
num_key_value_heads: Optional[int] = None
v_head_dim: Optional[int] = None
vocab_size: Optional[int] = None
kv_lora_rank: Optional[int] = None
qk_rope_head_dim: Optional[int] = None
qk_nope_head_dim: Optional[int] = None
# DSv4-specific (DSv4-Pro: per-layer compression ratios + sparse indexer + SWA)
compression_ratios: Optional[list] = None # per-layer: 4 or 128
indexer_head_dim: Optional[int] = None
window_size: Optional[int] = None
torch_dtype: Optional[str] = None
# deepseek v4 model config
qk_nope_head_dim: Optional[int] = None
qk_rope_head_dim: Optional[int] = None
indexer_head_dim: Optional[int] = None
def is_mla(self) -> bool:
return self.attention_arch == "MLA"
def is_dsv4(self) -> bool:
return self.compression_ratios is not None
@@ -0,0 +1,19 @@
from sglang_simulator.time_predictor.aiconfigurator import (
AIConfiguratorTimePredictor,
)
from sglang_simulator.time_predictor.base import (
InferTimePredictor,
ScheduleBatch,
ScheduleRequest,
)
from sglang_simulator.time_predictor.ml import MLTimePredictor
from sglang_simulator.time_predictor.replay import ReplayTimePredictor
__all__ = (
ScheduleRequest,
ScheduleBatch,
InferTimePredictor,
AIConfiguratorTimePredictor,
MLTimePredictor,
ReplayTimePredictor,
)
@@ -0,0 +1,291 @@
from typing import Optional
import numpy as np
from aiconfigurator.sdk import models
from aiconfigurator.sdk.backends.factory import get_backend
from aiconfigurator.sdk.common import (
CommQuantMode,
DatabaseMode,
FMHAQuantMode,
GEMMQuantMode,
KVCacheQuantMode,
MoEQuantMode,
)
from aiconfigurator.sdk.config import ModelConfig, RuntimeConfig
from aiconfigurator.sdk.inference_session import InferenceSession
from aiconfigurator.sdk.perf_database import get_database, get_systems_paths
from sglang_simulator.simulation.types import (
SchedulerConfig,
)
from sglang_simulator.spec.accelerator import AcceleratorInfo
from sglang_simulator.spec.data_type import DataType
from sglang_simulator.spec.model import ModelInfo
from sglang_simulator.time_predictor.base import (
InferTimePredictor,
ScheduleBatch,
ScheduleRequest,
)
from sglang_simulator.utils import get_logger
# Map the common data types to AIConfigurator data types.
MAP_DTYPE_TO_GEMMQuantMode = {
DataType.FP16: GEMMQuantMode.bfloat16,
DataType.BF16: GEMMQuantMode.bfloat16,
DataType.FP8: GEMMQuantMode.fp8_block,
DataType.INT8: GEMMQuantMode.int8_wo,
DataType.FP4: GEMMQuantMode.nvfp4,
DataType.INT4: GEMMQuantMode.int4_wo,
DataType.FP16_TENSOR: GEMMQuantMode.bfloat16,
DataType.BF16_TENSOR: GEMMQuantMode.bfloat16,
DataType.FP8_TENSOR: GEMMQuantMode.fp8,
DataType.INT8_TENSOR: GEMMQuantMode.int8_wo,
DataType.FP4_TENSOR: GEMMQuantMode.nvfp4,
DataType.INT4_TENSOR: GEMMQuantMode.int4_wo,
}
MAP_DTYPE_TO_KVCacheQuantMode = {
DataType.FP16: KVCacheQuantMode.bfloat16,
DataType.BF16: KVCacheQuantMode.bfloat16,
DataType.FP8: KVCacheQuantMode.fp8,
DataType.INT8: KVCacheQuantMode.int8,
}
MAP_DTYPE_TO_FMHAQuantMode = {
DataType.FP16: FMHAQuantMode.bfloat16,
DataType.BF16: FMHAQuantMode.bfloat16,
DataType.FP8: FMHAQuantMode.fp8,
}
MAP_DTYPE_TO_MoEQuantMode = {
DataType.FP16: MoEQuantMode.bfloat16,
DataType.BF16: MoEQuantMode.bfloat16,
DataType.FP8: MoEQuantMode.fp8_block,
DataType.INT8: MoEQuantMode.fp8,
DataType.FP4: MoEQuantMode.nvfp4,
DataType.INT4: MoEQuantMode.int4_wo,
}
MAP_DTYPE_TO_CommQuantMode = {
DataType.FP16: CommQuantMode.half,
DataType.BF16: CommQuantMode.half,
DataType.FP8: CommQuantMode.fp8,
DataType.INT8: CommQuantMode.int8,
}
logger = get_logger("sgl_simulator")
def _resolve_comm_quant_mode(sched_config: SchedulerConfig) -> CommQuantMode:
if sched_config.comm_quant_mode_override:
return getattr(CommQuantMode, sched_config.comm_quant_mode_override)
if sched_config.data_type is None:
return CommQuantMode.half
try:
return MAP_DTYPE_TO_CommQuantMode[sched_config.data_type]
except KeyError:
raise ValueError(
"AIConfigurator has no communication quantization mapping for "
f"model data type {sched_config.data_type.value}. Set "
"comm_quant_mode_override explicitly to half, int8, or fp8."
) from None
def get_perf_model(
sched_config: SchedulerConfig,
model: ModelInfo,
workload_distribution: str = "balanced",
) -> models.BaseModel:
model_config = ModelConfig(
pp_size=sched_config.pp_size,
tp_size=sched_config.attn_tp_size,
moe_tp_size=sched_config.moe_tp_size,
moe_ep_size=sched_config.moe_ep_size,
attention_dp_size=sched_config.attn_dp_size,
cp_size=sched_config.cp_size,
cp_style=sched_config.cp_style,
gemm_quant_mode=MAP_DTYPE_TO_GEMMQuantMode.get(
sched_config.data_type, GEMMQuantMode.bfloat16
),
moe_quant_mode=(
getattr(MoEQuantMode, sched_config.moe_quant_mode_override)
if sched_config.moe_quant_mode_override
else MAP_DTYPE_TO_MoEQuantMode.get(
sched_config.data_type, MoEQuantMode.bfloat16
)
),
kvcache_quant_mode=MAP_DTYPE_TO_KVCacheQuantMode.get(
sched_config.kv_cache_data_type, KVCacheQuantMode.bfloat16
),
fmha_quant_mode=(
getattr(FMHAQuantMode, sched_config.fmha_quant_mode_override)
if sched_config.fmha_quant_mode_override
else MAP_DTYPE_TO_FMHAQuantMode.get(
sched_config.kv_cache_data_type, FMHAQuantMode.bfloat16
)
),
comm_quant_mode=_resolve_comm_quant_mode(sched_config),
workload_distribution=workload_distribution,
)
logger.info(f"Model config for AIConfigurator: {model_config}")
return models.get_model(
model_path=model.model_path,
model_config=model_config,
backend_name=sched_config.backend_name,
)
class AIConfiguratorTimePredictor(InferTimePredictor):
def __init__(
self,
model: ModelInfo,
hw: AcceleratorInfo,
config: SchedulerConfig,
database_path: Optional[str] = None,
database_mode: DatabaseMode | str = DatabaseMode.SILICON,
prefill_scale_factor: float = 1,
decode_scale_factor: float = 1,
prefill_min_latency: float = 0,
workload_distribution: str = "balanced",
enable_oom_check: bool = False,
):
super().__init__(model, hw, config)
self.prefill_scale_factor = prefill_scale_factor
self.decode_scale_factor = decode_scale_factor
self.prefill_min_latency = prefill_min_latency
if isinstance(database_mode, str):
database_mode = self._get_database_mode(database_mode)
database = get_database(
system=hw.name,
backend=config.backend_name,
version=config.backend_version,
systems_paths=(
[database_path] if database_path is not None else get_systems_paths()
),
)
if database is None:
raise ValueError("Failed to initialize the database.")
database.set_default_database_mode(database_mode)
logger.info(f"AIC Database mode: {database_mode}")
self._session = InferenceSession(
model=get_perf_model(config, model, workload_distribution),
backend=get_backend(self.config.backend_name),
database=database,
)
self.enable_oom_check = enable_oom_check
self._is_oom = False
def _get_database_mode(self, mode: str) -> DatabaseMode:
return {
"SILICON": DatabaseMode.SILICON,
"HYBRID": DatabaseMode.HYBRID,
"EMPIRICAL": DatabaseMode.EMPIRICAL,
"SOL": DatabaseMode.SOL,
"SOL_FULL": DatabaseMode.SOL_FULL,
}.get(mode.upper(), DatabaseMode.SILICON)
def ctx_attn_flops_ratio_with_avg(self, reqs: list[ScheduleRequest]) -> float:
if len(reqs) == 1:
return 1.0
mean_past = np.mean([req.past_kv_length for req in reqs])
mean_input = np.mean([req.extend_length for req in reqs])
avg_flops = (mean_past + mean_past + mean_input) * mean_input / 2 * len(reqs)
actual_flops = 0
for req in reqs:
actual_flops += (
(req.past_kv_length + req.past_kv_length + req.extend_length)
* req.extend_length
/ 2
)
return actual_flops / avg_flops
def predict_infer_latency_dict(self, batch: ScheduleBatch) -> dict:
# Returns latency details for debugging operators.
if batch.is_decode():
# Decode: output sequence length (osl) = 2, input sequence length (isl) = mean(past_kv_length)
isl = int(np.mean([req.past_kv_length for req in batch.reqs]))
runtime_config = RuntimeConfig(batch_size=batch.batch_size, isl=isl, osl=2)
if self.enable_oom_check:
summary = self._session.run_static(runtime_config, mode="static_gen")
latency_dict = summary.get_generation_latency_dict()
else:
# faster path
results = self._session._backend._run_static_breakdown(
self._session._model,
self._session._database,
runtime_config,
mode="static_gen",
)
latency_dict = results[2]
else:
# Prefill: output sequence length (osl) = 1, input sequence length (isl) = mean(past_kv + input), prefix = mean(past_kv)
mean_past = np.mean([req.past_kv_length for req in batch.reqs])
mean_input = np.mean([req.extend_length for req in batch.reqs])
isl = int(mean_past + mean_input)
prefix = int(mean_past)
runtime_config = RuntimeConfig(
batch_size=batch.batch_size, isl=isl, prefix=prefix, osl=1
)
seq_imbalance_correction_scale = self.ctx_attn_flops_ratio_with_avg(
batch.reqs
)
if seq_imbalance_correction_scale >= 0.4:
runtime_config = RuntimeConfig(
batch_size=batch.batch_size,
isl=isl,
prefix=prefix,
osl=1,
seq_imbalance_correction_scale=seq_imbalance_correction_scale,
)
else:
runtime_config = RuntimeConfig(
batch_size=batch.batch_size, isl=isl, prefix=prefix, osl=1
)
if self.enable_oom_check:
summary = self._session.run_static(runtime_config, mode="static_ctx")
latency_dict = summary.get_context_latency_dict()
else:
# faster path
results = self._session._backend._run_static_breakdown(
self._session._model,
self._session._database,
runtime_config,
mode="static_ctx",
)
latency_dict = results[0]
return latency_dict
def predict_infer_time(self, batch: ScheduleBatch) -> float:
latency_dict = self.predict_infer_latency_dict(batch)
infer_time = sum(latency_dict.values())
if self._is_oom:
logger.warning("Out of memory detected during estimation.")
infer_time = -infer_time
if batch.is_decode():
infer_time *= self.decode_scale_factor
else:
infer_time *= self.prefill_scale_factor
if not batch.is_decode():
infer_time = (
max(infer_time, self.prefill_min_latency)
if infer_time > 0
else infer_time
)
return infer_time / 1e3
@@ -0,0 +1,97 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from sglang_simulator.simulation.types import SchedulerConfig
from sglang_simulator.spec.accelerator import AcceleratorInfo
from sglang_simulator.spec.model import ModelInfo
from sglang_simulator.utils import get_logger
logger = get_logger("sgl_simulator")
@dataclass
class ScheduleRequest:
extend_length: int = 0
past_kv_length: int = 0
@dataclass
class ScheduleBatch:
reqs: list[ScheduleRequest] = field(default_factory=list)
def __repr__(self) -> str:
return f"batch_size={len(self.reqs)},reqs={[(req.extend_length, req.past_kv_length) for req in self.reqs]}"
def __eq__(self, batch: "ScheduleBatch"):
if self.batch_size != batch.batch_size:
return False
req1, req2 = [], []
for idx in range(self.batch_size):
req1.append((self.reqs[idx].extend_length, self.reqs[idx].past_kv_length))
req2.append((batch.reqs[idx].extend_length, batch.reqs[idx].past_kv_length))
return sorted(req1) == sorted(req2)
def request_info(self) -> list[list[int, int]]:
# The request information organized in the format `(input_len, past_kv_len)`
return [[req.extend_length, req.past_kv_length] for req in self.reqs]
@property
def num_context_tokens(self) -> int:
return sum(req.extend_length for req in self.reqs)
@property
def total_past_kv_length(self) -> int:
return sum(req.past_kv_length for req in self.reqs)
@property
def batch_size(self) -> int:
return len(self.reqs)
def is_empty(self) -> bool:
return len(self.reqs) == 0
def is_prefill(self) -> bool:
return not self.is_decode()
def is_decode(self) -> bool:
for req in self.reqs:
if req.extend_length > 1:
return False
return True
@property
def num_ctx_requests(self) -> int:
return self.batch_size if self.is_prefill() else 0
@property
def num_gen_requests(self) -> int:
return self.batch_size if self.is_decode() else 0
class InferTimePredictor(ABC):
def __init__(
self,
model: ModelInfo,
hw: AcceleratorInfo,
config: SchedulerConfig,
*args,
**kwargs,
):
self.model: ModelInfo = model
self.hw: AcceleratorInfo = hw
self.config: SchedulerConfig = config
@abstractmethod
def predict_infer_time(self, batch: ScheduleBatch) -> float:
# Return the inference time in seconds. Return a negative value if an exception occurs (e.g., out of memory).
pass
def get_metrics(self) -> dict:
"""Return predictor-specific metrics for the current profile interval."""
return {}
def reset_metrics(self) -> None:
"""Reset predictor-specific metrics after a profile flush."""
return None
@@ -0,0 +1,156 @@
"""ML-trained per-iter latency predictor.
Loads a joblib pickle of a sklearn-compatible regressor and predicts forward latency
from batch composition features. Train one with `train_latency_model.py`.
sim_config.json usage:
"predictor": {
"name": "ml",
"database_path": "/path/to/latency_model.pkl"
}
"""
import math
import os
import joblib
from sglang_simulator.simulation.types import SchedulerConfig
from sglang_simulator.spec.accelerator import AcceleratorInfo
from sglang_simulator.spec.model import ModelInfo
from sglang_simulator.time_predictor.base import InferTimePredictor, ScheduleBatch
from sglang_simulator.utils import get_logger
logger = get_logger("sgl_simulator")
class MLTimePredictor(InferTimePredictor):
"""Per-iter latency predictor backed by an offline-trained sklearn regressor.
Features (18 dim) extracted from ScheduleBatch:
batch_size, sum/max/min(extend), sum/max/min(past),
sum(extend*past), sum(extend^2), sum(past^2),
sum_attn_flops (= sum(e*(p+e/2))),
sum(extend × max_past), log1p(sum_past), log1p(sum_attn_flops),
batch_size × sum_extend, max_past - min_past,
is_decode, is_prefill
"""
# This ordered list is the ABI between offline training and simulation.
# The concrete regressor algorithm is intentionally unrestricted as long
# as it exposes sklearn-compatible predict([[18 features]]) -> [seconds].
FEATURE_NAMES = [
"batch_size",
"sum_extend",
"max_extend",
"min_extend",
"sum_past",
"max_past",
"min_past",
"sum_extend_x_past",
"sum_extend_squared",
"sum_past_squared",
"sum_attn_flops",
"sum_extend_x_max_past",
"log1p_sum_past",
"log1p_sum_attn_flops",
"batch_size_x_sum_extend",
"max_past_minus_min_past",
"is_decode",
"is_prefill",
]
def __init__(
self,
model: ModelInfo,
hw: AcceleratorInfo,
config: SchedulerConfig,
database_path: str,
latency_scale: float = 1.0,
**kwargs,
):
super().__init__(model, hw, config)
database_path = os.path.expandvars(os.path.expanduser(database_path))
if not database_path or not os.path.exists(database_path):
raise FileNotFoundError(
f"MLTimePredictor database_path not found: {database_path}. "
"Train one with `train_latency_model.py` first."
)
bundle = joblib.load(database_path)
if (
not isinstance(bundle, dict)
or "model" not in bundle
or "features" not in bundle
):
raise ValueError(
"MLTimePredictor requires a joblib bundle containing both "
"'model' and ordered 'features' metadata"
)
self._model = bundle["model"]
saved_features = list(bundle["features"])
if saved_features != self.FEATURE_NAMES:
raise ValueError(
"MLTimePredictor feature contract mismatch: "
f"saved={saved_features}, expected={self.FEATURE_NAMES}. "
"Retrain or export the model with the exact 18-feature ABI."
)
if not callable(getattr(self._model, "predict", None)):
raise TypeError(
"MLTimePredictor model must expose a callable predict() method"
)
self._features = saved_features
self._call_count = 0
self._latency_scale = float(latency_scale)
logger.info(
"MLTimePredictor loaded from %s (model=%s, n_features=%d, latency_scale=%.4f)",
database_path,
type(self._model).__name__,
len(self._features),
self._latency_scale,
)
def predict_infer_time(self, batch: ScheduleBatch) -> float:
if batch.is_empty():
return 0.0
exts = [req.extend_length for req in batch.reqs]
pasts = [req.past_kv_length for req in batch.reqs]
bs = len(exts)
sum_e = sum(exts)
sum_p = sum(pasts)
sum_ep = sum(e * p for e, p in zip(exts, pasts))
sum_e2 = sum(e * e for e in exts)
sum_p2 = sum(p * p for p in pasts)
sum_attn = sum(e * (p + e / 2) for e, p in zip(exts, pasts))
max_e = max(exts)
max_p = max(pasts)
min_e = min(exts)
min_p = min(pasts)
feats = [
bs,
sum_e,
max_e,
min_e,
sum_p,
max_p,
min_p,
sum_ep,
sum_e2,
sum_p2,
sum_attn,
sum_e * max_p,
math.log1p(sum_p),
math.log1p(sum_attn),
bs * sum_e,
max_p - min_p,
int(all(e == 1 for e in exts)),
int(any(e > 1 for e in exts)),
]
self._call_count += 1
return float(self._model.predict([feats])[0]) * self._latency_scale
@@ -0,0 +1,189 @@
"""Oracle lookup predictor — replays real GPU iter_latency from a pre-built table.
Replay is a diagnostic predictor for separating latency-prediction error from
scheduler, cache, and simulator behavior. A replay table maps a JSON-encoded,
sorted list of ``[extend_input_length, prefix_length]`` pairs to measured
iteration latency in seconds.
Example simulator configuration:
"predictor": {
"name": "replay",
"database_path": "/path/to/replay_table.json",
"miss_strategy": "knn", # "zero" (legacy) or "knn" (interpolated)
"miss_knn_k": 3, # KNN k for "knn" strategy
"miss_fallback_seconds": 0.0 # used only when strategy=="zero"
}
"""
import json
import math
import os
from sglang_simulator.simulation.types import SchedulerConfig
from sglang_simulator.spec.accelerator import AcceleratorInfo
from sglang_simulator.spec.model import ModelInfo
from sglang_simulator.time_predictor.base import InferTimePredictor, ScheduleBatch
from sglang_simulator.utils import get_logger
logger = get_logger("sgl_simulator")
def _decode_key(key: str):
"""Parse a sorted-tuple lookup key back to list of (extend, past) pairs."""
return [tuple(pair) for pair in json.loads(key)]
def _shape_feat(extends, pasts):
"""3-D feature for KNN: (batch_size, sum_extend, sum_past).
Low-dim and aligned with the dominant latency drivers; avoids the curse of
dimensionality on the typically small (~1-3K) entries per replay table.
"""
return (len(extends), sum(extends), sum(pasts))
class ReplayTimePredictor(InferTimePredictor):
"""Oracle lookup predictor: returns real GPU iter_latency for matching batch compositions.
Compositions are matched exactly by sorted (extend_len, past_kv_len) tuples.
On lookup miss the behavior is controlled by `miss_strategy`:
- "zero" (default, legacy): returns `miss_fallback_seconds` (default 0.0).
Use for "is the gap NOT from the predictor?" diagnostic. Only meaningful
when miss rate is small AND you accept the bias of dropping miss work.
- "knn": KNN-interpolated latency from the k batches in the table with
the closest (batch_size, sum_extend, sum_past) shape (per-dim z-scored
Euclidean distance). Use when miss rate matters or you want a logically
complete oracle. Typical k=3.
For workloads where sim batch composition diverges from real (e.g., bursty
max-tps), the miss rate can be high — always check the post-run hit ratio.
"""
def __init__(
self,
model: ModelInfo,
hw: AcceleratorInfo,
config: SchedulerConfig,
database_path: str,
miss_fallback_seconds: float = 0.0,
miss_strategy: str = "zero",
miss_knn_k: int = 3,
**kwargs,
):
super().__init__(model, hw, config)
if not database_path or not os.path.exists(database_path):
raise FileNotFoundError(
f"ReplayTimePredictor database_path not found: {database_path}"
)
with open(database_path) as f:
self._table = json.load(f)
if miss_strategy not in ("zero", "knn"):
raise ValueError(
f"miss_strategy must be 'zero' or 'knn', got {miss_strategy!r}"
)
self._miss_strategy = miss_strategy
self._miss_fallback = float(miss_fallback_seconds)
self._miss_knn_k = int(miss_knn_k)
self._hits = 0
self._misses = 0
if self._miss_strategy == "knn":
self._prep_knn_index()
logger.info(
"ReplayTimePredictor loaded %d unique compositions from %s "
"(miss_strategy=knn, k=%d)",
len(self._table),
database_path,
self._miss_knn_k,
)
else:
self._knn_feats = None
logger.info(
"ReplayTimePredictor loaded %d unique compositions from %s "
"(miss_strategy=zero, fallback=%.4fs)",
len(self._table),
database_path,
self._miss_fallback,
)
def _prep_knn_index(self):
"""Build per-feature mean/std + cached (feat, lat) arrays for KNN fallback.
Uses plain Python (no numpy / sklearn dependency from the predictor side)
— table size is ~1K-3K so this is fine. Distance is z-scored Euclidean
across (batch_size, sum_extend, sum_past).
"""
feats = []
lats = []
for key, lat in self._table.items():
extends_pasts = _decode_key(key)
extends = [e for e, _ in extends_pasts]
pasts = [p for _, p in extends_pasts]
feats.append(_shape_feat(extends, pasts))
lats.append(float(lat))
n = len(feats)
# per-dim mean/std for z-scoring
means = [sum(f[d] for f in feats) / n for d in range(3)]
var = [sum((f[d] - means[d]) ** 2 for f in feats) / n for d in range(3)]
stds = [math.sqrt(v) if v > 1e-12 else 1.0 for v in var]
self._knn_feats = feats
self._knn_lats = lats
self._knn_mean = means
self._knn_std = stds
def _knn_predict(self, query_feat):
"""k nearest neighbors over z-scored 3-D shape feature, simple mean."""
qz = tuple(
(query_feat[d] - self._knn_mean[d]) / self._knn_std[d] for d in range(3)
)
# compute squared distance to every table entry; pick smallest k
# n ≤ ~3K so O(n) per query is fine; sim only calls on misses
dists = []
for i, f in enumerate(self._knn_feats):
fz = (
(f[0] - self._knn_mean[0]) / self._knn_std[0],
(f[1] - self._knn_mean[1]) / self._knn_std[1],
(f[2] - self._knn_mean[2]) / self._knn_std[2],
)
d2 = (fz[0] - qz[0]) ** 2 + (fz[1] - qz[1]) ** 2 + (fz[2] - qz[2]) ** 2
dists.append((d2, i))
dists.sort(key=lambda t: t[0])
k = min(self._miss_knn_k, len(dists))
return sum(self._knn_lats[i] for _, i in dists[:k]) / k
def predict_infer_time(self, batch: ScheduleBatch) -> float:
if batch.is_empty():
return 0.0
extends = [req.extend_length for req in batch.reqs]
pasts = [req.past_kv_length for req in batch.reqs]
key = json.dumps(
sorted([req.extend_length, req.past_kv_length] for req in batch.reqs)
)
v = self._table.get(key)
if v is None:
self._misses += 1
if self._miss_strategy == "knn":
return self._knn_predict(_shape_feat(extends, pasts))
return self._miss_fallback
self._hits += 1
return float(v)
def get_metrics(self) -> dict:
total = self._hits + self._misses
return {
"replay_exact_match_steps": self._hits,
"replay_miss_steps": self._misses,
"replay_zero_fallback_steps": (
self._misses if self._miss_strategy == "zero" else 0
),
"replay_knn_fallback_steps": (
self._misses if self._miss_strategy == "knn" else 0
),
"replay_fallback_rate": self._misses / total if total else 0.0,
}
def reset_metrics(self) -> None:
self._hits = 0
self._misses = 0
@@ -0,0 +1,3 @@
from sglang_simulator.utils.logger import get_logger
__all__ = ["get_logger"]
@@ -0,0 +1,22 @@
import json
from dataclasses import asdict, is_dataclass
from enum import Enum
import numpy as np
class CustomJsonEncoder(json.JSONEncoder):
def default(self, obj):
# Enum
if isinstance(obj, Enum):
return obj.value
# Dataclass
if is_dataclass(obj):
return asdict(obj)
# Numpy
if isinstance(obj, (np.int32, np.int64, np.float32, np.float64)):
return int(obj) if isinstance(obj, (np.int32, np.int64)) else float(obj)
if isinstance(obj, np.ndarray):
return obj.tolist()
# Other
return super().default(obj)
@@ -0,0 +1,16 @@
import logging
def get_logger(name: str = "sglang_simulator") -> logging.Logger:
logger = logging.getLogger(name)
if not logger.handlers:
logger.setLevel(logging.INFO)
handler = logging.StreamHandler()
formatter = logging.Formatter(
fmt="%(asctime)s %(levelname)s [%(name)s] %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
handler.setFormatter(formatter)
logger.addHandler(handler)
logger.propagate = False
return logger
@@ -0,0 +1,15 @@
"""Early CPU-simulation compatibility for spawned SGLang workers."""
import os
if (
os.environ.get("SGLANG_SIMULATOR_BOOTSTRAP") == "1"
and os.environ.get("SGLANG_USE_CPU_ENGINE") == "1"
):
import torch
# Some model-specific import-time checks probe the target GPU even though
# SGLang Simulator never executes a real model forward. Spawned workers reach those
# imports before the simulator target wrapper can run. CPU simulation is
# explicit, so physical GPU visibility must not affect this shim.
torch.cuda.get_device_capability = lambda *_args, **_kwargs: (10, 0)
@@ -0,0 +1,28 @@
{
"architectures": ["Qwen3ForCausalLM"],
"attention_bias": false,
"attention_dropout": 0.0,
"bos_token_id": 151643,
"eos_token_id": 151645,
"head_dim": 128,
"hidden_act": "silu",
"hidden_size": 4096,
"initializer_range": 0.02,
"intermediate_size": 12288,
"max_position_embeddings": 40960,
"max_window_layers": 36,
"model_type": "qwen3",
"num_attention_heads": 32,
"num_hidden_layers": 36,
"num_key_value_heads": 8,
"rms_norm_eps": 1e-06,
"rope_scaling": null,
"rope_theta": 1000000,
"sliding_window": null,
"tie_word_embeddings": false,
"torch_dtype": "bfloat16",
"transformers_version": "5.12.1",
"use_cache": true,
"use_sliding_window": false,
"vocab_size": 151936
}
@@ -0,0 +1,68 @@
import json
import pytest
from sglang_simulator.simulation.benchmark import BenchmarkConfig
from test_simulation_sglang_runner import make_fixed_dataset, make_sglang_runner
from test_simulation_sglang_serving import (
SIM_CONFIGS,
SGLangServingRunner,
assert_decode_metrics,
)
def test_in_process_runner_reports_each_cache_tier(tmp_path):
runner = make_sglang_runner(tmp_path)
benchmark_config = BenchmarkConfig(request_rate=10, ignore_request_timestamp=False)
cached_ds = make_fixed_dataset(1000, 8)
evict_l1_ds = make_fixed_dataset(2000, 10)
evict_l2_ds = make_fixed_dataset(3000, 20)
try:
metrics = runner.benchmark(benchmark_config, dataset=cached_ds)
assert metrics["completed"] == len(cached_ds)
assert metrics["prefix_cache_reused_ratio"] == 0
metrics = runner.benchmark(benchmark_config, dataset=cached_ds)
assert metrics["kv_cache_device_hit_ratio"] > 0.95
runner.benchmark(benchmark_config, dataset=evict_l1_ds)
metrics = runner.benchmark(benchmark_config, dataset=cached_ds)
assert metrics["kv_cache_host_hit_ratio"] > 0.95
runner.benchmark(benchmark_config, dataset=evict_l2_ds)
metrics = runner.benchmark(benchmark_config, dataset=cached_ds)
assert metrics["kv_cache_storage_hit_ratio"] > 0.95
finally:
runner.shutdown()
def test_second_replay_benchmark_hits_all_reusable_prefix_tokens(tmp_path, monkeypatch):
# This test validates cache reuse across consecutive benchmark runs.
monkeypatch.setenv("SGLANG_IS_IN_CI", "false")
runner = SGLangServingRunner(SIM_CONFIGS["replay"], tmp_path)
try:
first_metrics = runner.benchmark(tmp_path / "benchmark-first.json")
second_metrics = runner.benchmark(tmp_path / "benchmark-second.json")
finally:
runner.shutdown()
assert_decode_metrics(first_metrics)
assert_decode_metrics(second_metrics)
assert second_metrics["total_input"] == 24
assert second_metrics["total_new_input"] == 3
assert second_metrics["prefix_cache_reused_ratio"] == pytest.approx(0.875)
assert second_metrics["kv_cache_device_hit_ratio"] == pytest.approx(0.875)
assert second_metrics["kv_cache_host_hit_ratio"] == 0
assert second_metrics["kv_cache_storage_hit_ratio"] == 0
requests = [
json.loads(line)
for line in (runner.output_dir / "request.jsonl")
.read_text(encoding="utf-8")
.splitlines()
]
assert len(requests) == 3
assert all(request["input_length"] == 8 for request in requests)
assert all(request["final_device_hit_len"] == 7 for request in requests)
@@ -0,0 +1,80 @@
import json
import pytest
from test_simulation_sglang_serving import (
SIM_CONFIGS,
SGLangServingRunner,
assert_decode_metrics,
)
REQUEST_RATE = 1
SEED = 123
RELATIVE_TOLERANCES = {
"duration": 0.01,
"request_throughput": 0.01,
"input_throughput": 0.01,
"output_throughput": 0.01,
"mean_e2e_latency_ms": 0.10,
"mean_ttft_ms": 0.10,
"mean_tpot_ms": 0.10,
"mean_itl_ms": 0.10,
}
def _relative_error(actual, expected):
return abs(actual - expected) / abs(expected)
def _run_mode(mode, tmp_path):
case_dir = tmp_path / mode
case_dir.mkdir()
runner = SGLangServingRunner(SIM_CONFIGS["aic_sol"], case_dir, mode=mode)
try:
metrics = runner.benchmark(
case_dir / "benchmark.json", request_rate=REQUEST_RATE, seed=SEED
)
finally:
runner.shutdown()
requests = [
json.loads(line)
for line in (runner.output_dir / "request.jsonl")
.read_text(encoding="utf-8")
.splitlines()
]
requests.sort(key=lambda request: request["created_time"])
return metrics, requests
def test_request_rate_offline_matches_blocking(tmp_path):
offline_metrics, offline_requests = _run_mode("offline", tmp_path)
blocking_metrics, blocking_requests = _run_mode("blocking", tmp_path)
for metrics in (offline_metrics, blocking_metrics):
assert_decode_metrics(metrics)
assert len(offline_requests) == len(blocking_requests) == 3
offline_arrivals = [request["created_time"] for request in offline_requests]
blocking_arrivals = [request["created_time"] for request in blocking_requests]
assert offline_arrivals[1] > 0.5
assert blocking_arrivals[1] > 0.5
assert offline_arrivals == pytest.approx(blocking_arrivals, abs=0.02)
assert (
offline_metrics["max_concurrent_requests"]
== blocking_metrics["max_concurrent_requests"]
== 1
)
for key in ("completed", "total_input", "total_output"):
assert offline_metrics[key] == blocking_metrics[key]
for key, tolerance in RELATIVE_TOLERANCES.items():
error = _relative_error(offline_metrics[key], blocking_metrics[key])
assert error <= tolerance, (
key,
offline_metrics[key],
blocking_metrics[key],
error,
tolerance,
)
@@ -0,0 +1,118 @@
import atexit
import json
import os
import sys
from pathlib import Path
from unittest.mock import patch
from sglang_simulator.dataset import GenericRequest, SimpleDataset
from sglang_simulator.simulation.benchmark import BenchmarkConfig
ASSETS = Path(__file__).parent / "assets"
SGLANG_ROOT = Path(__file__).parents[3]
if str(SGLANG_ROOT) not in sys.path:
sys.path.insert(0, str(SGLANG_ROOT))
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "")
def make_fixed_dataset(
start_token: int,
count: int,
*,
input_length: int = 1025,
output_length: int = 1,
) -> SimpleDataset:
return SimpleDataset(
reqs=[
GenericRequest(
token_ids=[start_token + i] * input_length,
input_length=input_length,
output_length=output_length,
custom_params={"created_time": i / 10},
)
for i in range(count)
]
)
def _write_sim_config(tmp_path: Path) -> Path:
table_path = tmp_path / "replay.json"
table_path.write_text(
json.dumps({"[[1, 1024]]": 0.001, "[[1025, 0]]": 0.01}),
encoding="utf-8",
)
config = {
"platform": {
"accelerator": {"name": "a100_sxm", "hbm_capacity_gb": 80},
"disk_read_bandwidth_gb": 8,
"disk_write_bandwidth_gb": 8,
"memory_read_bandwidth_gb": 64,
"memory_write_bandwidth_gb": 64,
"num_device_per_node": 8,
},
"predictor": {
"name": "replay",
"database_path": str(table_path),
"miss_strategy": "knn",
"miss_knn_k": 1,
},
"scheduler": {"tp_size": 1, "ep_size": 1, "dp_size": 1},
}
config_path = tmp_path / "sim_config.json"
config_path.write_text(json.dumps(config), encoding="utf-8")
return config_path
def make_sglang_runner(tmp_path: Path):
os.environ["SGLANG_SIMULATOR_CONFIG_PATH"] = str(_write_sim_config(tmp_path))
from benchmark.simulator.bench_runner import SGLangBenchmarkRunner
from sglang.srt.server_args import ServerArgs
runner = SGLangBenchmarkRunner(
server_args=ServerArgs(
model_path=str(ASSETS / "qwen3-8b"),
load_format="dummy",
device="cpu",
enable_hierarchical_cache=True,
hicache_ratio=2,
hicache_storage_backend="file",
hicache_storage_prefetch_policy="wait_complete",
max_total_tokens=10 * 1024,
page_size=256,
skip_tokenizer_init=True,
)
)
runner.clear_hicache_storage()
return runner
def test_benchmark_sglang_runs_paged_decode(tmp_path):
runner = make_sglang_runner(tmp_path)
dataset = make_fixed_dataset(
1000,
2,
input_length=1024,
output_length=2,
)
try:
metrics = runner.benchmark(
BenchmarkConfig(request_rate=10, ignore_request_timestamp=False),
dataset=dataset,
)
request_stats = runner.get_request_stats()
finally:
with patch.object(atexit, "unregister", wraps=atexit.unregister) as unregister:
runner.shutdown()
runner.shutdown()
unregister.assert_called_once_with(runner.engine.shutdown)
assert metrics["completed"] == len(dataset)
assert metrics["total_input"] == 2 * 1024
assert metrics["total_output"] == 2 * 2
assert metrics["mean_tpot_ms"] > 0
assert all(
idx == 0 or req["created_time"] > 0 for idx, req in enumerate(request_stats)
)
@@ -0,0 +1,163 @@
import json
import os
import signal
import socket
import subprocess
import sys
import time
from pathlib import Path
import pytest
import requests
ASSETS = Path(__file__).parent / "assets"
SGLANG_ROOT = Path(__file__).parents[3]
BENCH_SERVING = SGLANG_ROOT / "benchmark" / "simulator" / "bench_serving.py"
EXAMPLES = Path(__file__).parent.parent / "examples"
SIM_CONFIGS = {
"aic_sol": EXAMPLES / "sim_configs" / "aic_sol.json",
"aic_silicon": EXAMPLES / "sim_configs" / "aic_silicon.json",
"ml": EXAMPLES / "sim_configs" / "ml.json",
"replay": EXAMPLES / "sim_configs" / "replay.json",
}
class SGLangServingRunner:
def __init__(self, config_path: Path, tmp_path: Path, mode: str = "offline"):
self.mode = mode
with socket.socket() as sock:
sock.bind(("127.0.0.1", 0))
self.port = sock.getsockname()[1]
self.output_dir = tmp_path / "output"
env = os.environ.copy()
env.update(
CUDA_VISIBLE_DEVICES="",
SGLANG_USE_CPU_ENGINE="1",
SGLANG_SIMULATOR_CONFIG_PATH=str(config_path),
SGLANG_SIMULATOR_OUTPUT_MODE=mode.upper(),
SGLANG_SIMULATOR_OUTPUT_DIR=str(self.output_dir),
)
cmd = [
sys.executable,
"-m",
"sglang_simulator.simulation.sglang.launch_server",
"--model-path",
str(ASSETS / "qwen3-8b"),
"--sim-config-path",
str(config_path),
"--port",
str(self.port),
"--tokenizer-path",
str(EXAMPLES / "assets" / "tokenizer"),
"--max-total-tokens",
"8192",
"--max-running-requests",
"8",
"--disable-overlap-schedule",
]
self.server_proc = subprocess.Popen(cmd, env=env, preexec_fn=os.setsid)
for _ in range(120):
if self.server_proc.poll() is not None:
raise RuntimeError("SGLang Simulator server exited during startup")
try:
if requests.get(self.base_url, timeout=1).status_code < 500:
return
except requests.RequestException:
pass
time.sleep(1)
self.shutdown()
raise RuntimeError("SGLang Simulator server did not become ready")
@property
def base_url(self) -> str:
return f"http://127.0.0.1:{self.port}"
def benchmark(
self,
output_file: Path,
workload: str = "sharegpt",
request_rate=None,
seed=42,
) -> dict:
cmd = [
sys.executable,
str(BENCH_SERVING),
f"--simulator-mode={self.mode}",
"--backend=sglang",
f"--base-url={self.base_url}",
f"--model={ASSETS / 'qwen3-8b'}",
f"--tokenizer={EXAMPLES / 'assets' / 'tokenizer'}",
"--num-prompts=3",
"--disable-tqdm",
"--profile",
f"--output-file={output_file}",
]
if request_rate is not None:
cmd.extend([f"--request-rate={request_rate}", f"--seed={seed}"])
if workload == "sharegpt":
cmd.extend(
[
"--dataset-name=sharegpt",
f"--dataset-path={EXAMPLES / 'workloads' / 'sharegpt-example.json'}",
"--sharegpt-output-len=4",
]
)
else:
assert workload == "timestamp_trace"
cmd.extend(
[
"--dataset-name=autobench",
f"--dataset-path={EXAMPLES / 'workloads' / 'timestamp-trace-example.jsonl'}",
"--use-trace-timestamps",
]
)
subprocess.run(cmd, check=True)
assert output_file.is_file()
return json.loads(
(self.output_dir / "metrics.json").read_text(encoding="utf-8")
)
def shutdown(self):
if self.server_proc.poll() is not None:
return
os.killpg(self.server_proc.pid, signal.SIGTERM)
try:
self.server_proc.wait(timeout=10)
except subprocess.TimeoutExpired:
os.killpg(self.server_proc.pid, signal.SIGKILL)
self.server_proc.wait()
def assert_decode_metrics(metrics):
assert metrics["completed"] == 3
assert metrics["total_output"] == 12
assert metrics["mean_ttft_ms"] >= 0
assert metrics["mean_tpot_ms"] > 0
assert metrics["mean_itl_ms"] > 0
assert metrics["input_throughput"] > 0
@pytest.mark.parametrize("config_name", SIM_CONFIGS)
def test_benchmark(config_name, tmp_path):
runner = SGLangServingRunner(SIM_CONFIGS[config_name], tmp_path)
try:
metrics = runner.benchmark(tmp_path / "benchmark.json")
finally:
runner.shutdown()
assert_decode_metrics(metrics)
def test_timestamp_trace(tmp_path):
runner = SGLangServingRunner(SIM_CONFIGS["replay"], tmp_path)
try:
metrics = runner.benchmark(
tmp_path / "benchmark.json", workload="timestamp_trace"
)
finally:
runner.shutdown()
assert_decode_metrics(metrics)