[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:
co-authored by
zhouhaizhu.zhz
LinSiyuan814
hzh0425
parent
a5f07b1241
commit
59799a3687
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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/
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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",
|
||||
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
@@ -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}
|
||||
@@ -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"
|
||||
]
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user