config: the runtime readers take the published bags (#36254)
This commit is contained in:
@@ -30,6 +30,7 @@ from sglang.benchmark.datasets import DatasetRow, get_dataset
|
||||
from sglang.benchmark.datasets.random import sample_random_requests
|
||||
from sglang.benchmark.utils import get_tokenizer, set_ulimit
|
||||
from sglang.lang.backend.runtime_endpoint import Runtime
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.entrypoints.engine import Engine
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
@@ -366,6 +367,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
|
||||
RayEngine requires a placement group, so we launch it inside a Ray actor
|
||||
and return a lightweight proxy that forwards calls via ray.get().
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
import ray
|
||||
from ray.runtime_env import RuntimeEnv
|
||||
from ray.util.placement_group import placement_group
|
||||
@@ -377,7 +379,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
|
||||
if not ray.is_initialized():
|
||||
ray.init(runtime_env=RuntimeEnv(env_vars=env_vars))
|
||||
|
||||
total_gpus = server_args.tp_size * server_args.pp_size
|
||||
total_gpus = cfg.tp_size * cfg.pp_size
|
||||
pg = placement_group([{"CPU": 1, "GPU": total_gpus}], strategy="STRICT_PACK")
|
||||
ray.get(pg.ready())
|
||||
|
||||
@@ -398,7 +400,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
|
||||
placement_group=pg,
|
||||
placement_group_bundle_index=0,
|
||||
),
|
||||
).remote(**dict(server_args._raw_input))
|
||||
).remote(**dict(cfg._raw_input))
|
||||
|
||||
class _Proxy:
|
||||
"""Forwards method calls to the remote RayEngine actor."""
|
||||
@@ -434,20 +436,21 @@ def throughput_test(
|
||||
):
|
||||
# A programmatic caller may hand over a freshly constructed record, and
|
||||
# the backends below read the resolved paths and the raw snapshot.
|
||||
server_args.resolve_once()
|
||||
cfg = resolving_view(server_args)
|
||||
cfg.resolve_once()
|
||||
if bench_args.backend == "engine":
|
||||
if server_args.use_ray:
|
||||
if cfg.use_ray:
|
||||
backend = _create_ray_engine_backend(server_args)
|
||||
else:
|
||||
backend = Engine(server_args=server_args)
|
||||
if not backend:
|
||||
raise ValueError("Please provide valid engine arguments")
|
||||
elif bench_args.backend == "runtime":
|
||||
backend = Runtime(**dict(server_args._raw_input))
|
||||
backend = Runtime(**dict(cfg._raw_input))
|
||||
else:
|
||||
raise ValueError('Please set backend to either "engine" or "runtime"')
|
||||
|
||||
tokenizer_id = server_args.tokenizer_path or server_args.model_path
|
||||
tokenizer_id = cfg.tokenizer_path or cfg.model_path
|
||||
tokenizer = get_tokenizer(tokenizer_id)
|
||||
|
||||
# Set global environments
|
||||
|
||||
@@ -64,6 +64,7 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolution_result, resolving_view
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
@@ -79,7 +80,11 @@ from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
CudaGraphConfig,
|
||||
Phase,
|
||||
cuda_graph_fully_disabled,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule, publish
|
||||
@@ -297,44 +302,45 @@ class BenchArgs:
|
||||
|
||||
|
||||
def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
cfg = resolving_view(server_args)
|
||||
suppress_other_loggers()
|
||||
rank_print = print if tp_rank == 0 else lambda *args, **kwargs: None
|
||||
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
||||
moe_ep_rank = tp_rank // (cfg.tp_size // cfg.ep_size)
|
||||
|
||||
model_config = ModelConfig.from_server_args(server_args)
|
||||
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
|
||||
compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
cfg.enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
server_args.attn_cp_size,
|
||||
cfg.tp_size,
|
||||
cfg.dp_size,
|
||||
cfg.attn_cp_size,
|
||||
)
|
||||
)
|
||||
ps = ParallelState(
|
||||
tp_rank=tp_rank,
|
||||
tp_size=server_args.tp_size,
|
||||
tp_size=cfg.tp_size,
|
||||
pp_rank=0,
|
||||
pp_size=1,
|
||||
dp_rank=None,
|
||||
dp_size=server_args.dp_size,
|
||||
dp_size=cfg.dp_size,
|
||||
attn_tp_rank=attn_tp_rank,
|
||||
attn_tp_size=attn_tp_size,
|
||||
attn_cp_rank=0,
|
||||
attn_cp_size=server_args.attn_cp_size,
|
||||
attn_dcp_rank=tp_rank % server_args.dcp_size,
|
||||
attn_dcp_size=server_args.dcp_size,
|
||||
attn_cp_size=cfg.attn_cp_size,
|
||||
attn_dcp_rank=tp_rank % cfg.dcp_size,
|
||||
attn_dcp_size=cfg.dcp_size,
|
||||
attn_dp_rank=attn_dp_rank,
|
||||
attn_dp_size=attn_dp_size,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
moe_ep_size=server_args.ep_size,
|
||||
moe_ep_size=cfg.ep_size,
|
||||
moe_dp_rank=None,
|
||||
moe_dp_size=server_args.moe_dp_size,
|
||||
moe_dp_size=cfg.moe_dp_size,
|
||||
gpu_id=gpu_id,
|
||||
)
|
||||
runner_kwargs = dict(
|
||||
model_config=model_config,
|
||||
mem_fraction_static=server_args.mem_fraction_static,
|
||||
mem_fraction_static=cfg.mem_fraction_static,
|
||||
gpu_id=gpu_id,
|
||||
ps=ps,
|
||||
nccl_port=port_args.nccl_port,
|
||||
@@ -350,20 +356,20 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
model_runner = MlxModelRunnerStub(**runner_kwargs)
|
||||
else:
|
||||
model_runner = ModelRunner(**runner_kwargs)
|
||||
if server_args.is_startup_weight_load_overlap:
|
||||
if cfg.is_startup_weight_load_overlap:
|
||||
model_runner.start_startup_weight_load()
|
||||
model_runner.alloc_memory_pool()
|
||||
model_runner.init_attention_backends()
|
||||
model_runner.init_cuda_graphs()
|
||||
if server_args.is_startup_weight_load_overlap:
|
||||
if cfg.is_startup_weight_load_overlap:
|
||||
model_runner.finalize_startup_weight_load()
|
||||
rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}")
|
||||
tokenizer = get_tokenizer(
|
||||
server_args.tokenizer_path,
|
||||
tokenizer_mode=server_args.tokenizer_mode,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
cfg.tokenizer_path,
|
||||
tokenizer_mode=cfg.tokenizer_mode,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
)
|
||||
if server_args.tp_size > 1:
|
||||
if cfg.tp_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if _use_mlx:
|
||||
@@ -584,19 +590,20 @@ class _MlxBenchRunner:
|
||||
"""Wraps MlxModelRunner for the MLX benchmark path."""
|
||||
|
||||
def __init__(self, model_runner, server_args):
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||
|
||||
# Radix cache requires the scheduler's allocator/trie; disable in
|
||||
# standalone bench mode where no scheduler is present.
|
||||
init_kwargs = dict(
|
||||
model_path=server_args.model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
model_path=cfg.model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
disable_radix_cache=True,
|
||||
mem_fraction_static=server_args.mem_fraction_static,
|
||||
quantization=server_args.quantization,
|
||||
mem_fraction_static=cfg.mem_fraction_static,
|
||||
quantization=cfg.quantization,
|
||||
)
|
||||
if server_args.max_total_tokens is not None:
|
||||
init_kwargs["pool_size"] = server_args.max_total_tokens
|
||||
if cfg.max_total_tokens is not None:
|
||||
init_kwargs["pool_size"] = cfg.max_total_tokens
|
||||
self.mlx_runner = MlxModelRunner(**init_kwargs)
|
||||
self.mlx_runner.init_cache_pools(req_to_token_pool=None)
|
||||
self.fake_torch_runner = model_runner
|
||||
@@ -883,17 +890,18 @@ def latency_test(
|
||||
gpu_id,
|
||||
tp_rank,
|
||||
):
|
||||
cfg = resolving_view(server_args)
|
||||
# `main` runs this inline for tp_size == 1 and spawns it per rank otherwise;
|
||||
# a spawned child arrives with nothing published.
|
||||
publish(server_args, role="scheduler")
|
||||
initialize_moe_config(server_args)
|
||||
initialize_fp8_gemm_config(server_args)
|
||||
initialize_fp4_gemm_config(server_args)
|
||||
initialize_moe_config()
|
||||
initialize_fp8_gemm_config()
|
||||
initialize_fp4_gemm_config()
|
||||
|
||||
# Set CPU affinity
|
||||
if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"):
|
||||
parallel = get_parallel().config
|
||||
set_gpu_proc_affinity(
|
||||
server_args.pp_size, server_args.tp_size, server_args.nnodes, tp_rank
|
||||
parallel.pp_size, parallel.tp_size, parallel.nnodes, tp_rank
|
||||
)
|
||||
|
||||
# Configure the logger
|
||||
@@ -988,22 +996,45 @@ def latency_test(
|
||||
for result in result_list:
|
||||
fout.write(json.dumps(result) + "\n")
|
||||
|
||||
if server_args.tp_size > 1:
|
||||
if cfg.tp_size > 1:
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
def main(server_args, bench_args):
|
||||
# The decode phase has to capture the batch sizes this run benchmarks, and
|
||||
# the per-phase convenience knob loses to an explicit --cuda-graph-config
|
||||
# JSON (resolution applies that last), so the size is merged into that JSON.
|
||||
if getattr(server_args, "_declarations_materialized", False):
|
||||
# A record the caller already resolved: nothing will parse a raw dict
|
||||
# again, so the declaration has to be the finished typed config.
|
||||
merged = resolution_result(server_args, "cuda_graph_config")
|
||||
merged = (
|
||||
merged.to_dict()
|
||||
if isinstance(merged, CudaGraphConfig)
|
||||
else dict(merged or {})
|
||||
)
|
||||
decode = dict(merged.get(Phase.DECODE) or {})
|
||||
decode["max_bs"] = max(bench_args.batch_size)
|
||||
merged[Phase.DECODE] = decode
|
||||
graph_config = CudaGraphConfig.from_dict(merged)
|
||||
else:
|
||||
explicit = server_args.cuda_graph_config
|
||||
if isinstance(explicit, CudaGraphConfig):
|
||||
explicit = explicit.to_dict()
|
||||
graph_config = dict(explicit or {})
|
||||
decode = dict(graph_config.get(Phase.DECODE) or {})
|
||||
decode["max_bs"] = max(bench_args.batch_size)
|
||||
graph_config[Phase.DECODE] = decode
|
||||
server_args = server_args.replace_resolved(
|
||||
"benchmark.one_batch", cuda_graph_config=graph_config
|
||||
)
|
||||
server_args.resolve_once()
|
||||
|
||||
# The legacy cuda_graph_max_bs_decode field does not propagate; set the
|
||||
# decode phase.
|
||||
if server_args.cuda_graph_config is not None:
|
||||
server_args.cuda_graph_config[Phase.DECODE].max_bs = max(bench_args.batch_size)
|
||||
cfg = resolving_view(server_args)
|
||||
|
||||
_set_envs_and_config(server_args)
|
||||
|
||||
if server_args.model_path:
|
||||
if cfg.model_path:
|
||||
if bench_args.correctness_test:
|
||||
work_func = correctness_test
|
||||
else:
|
||||
|
||||
@@ -33,6 +33,7 @@ from sglang.benchmark.datasets import get_dataset
|
||||
from sglang.benchmark.endpoint import acquire_endpoint
|
||||
from sglang.benchmark.utils import get_processor, get_tokenizer
|
||||
from sglang.profiler import run_profile
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
||||
from sglang.srt.entrypoints.http_server import launch_server
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -1234,6 +1235,7 @@ def run_benchmark_internal(
|
||||
|
||||
|
||||
def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
||||
cfg = resolving_view(server_args)
|
||||
results, server_info = run_benchmark_internal(server_args, bench_args)
|
||||
|
||||
# Save results as pydantic models in the JSON format
|
||||
@@ -1241,7 +1243,7 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
||||
save_results_as_pydantic_models(
|
||||
results,
|
||||
pydantic_result_filename=bench_args.pydantic_result_filename,
|
||||
model_path=server_args.model_path,
|
||||
model_path=cfg.model_path,
|
||||
server_args=bench_args.server_args_for_metrics,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user