[Observability] Add startup, memory, and hybrid SWA diagnostics (#33375)

This commit is contained in:
Lianmin Zheng
2026-08-04 12:50:09 -07:00
committed by GitHub
parent 5081c063c0
commit 4794b401d5
36 changed files with 867 additions and 140 deletions
+1 -1
View File
@@ -691,7 +691,7 @@ class ModelConfig:
) )
if self.is_hybrid_swa: if self.is_hybrid_swa:
logger.info(f"Hybrid swa model: {self.hf_config.architectures=}") logger.debug(f"Hybrid swa model: {self.hf_config.architectures=}")
self.is_deepseek_v4_arch = any( self.is_deepseek_v4_arch = any(
arch arch
+35
View File
@@ -94,6 +94,7 @@ from sglang.srt.managers.multi_tokenizer_mixin import (
) )
from sglang.srt.managers.scheduler import run_scheduler_process from sglang.srt.managers.scheduler import run_scheduler_process
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.observability.startup_time import build_engine_startup_time
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.parser.template_detection import resolve_auto_parsers from sglang.srt.parser.template_detection import resolve_auto_parsers
from sglang.srt.parser.template_manager import TemplateManager from sglang.srt.parser.template_manager import TemplateManager
@@ -990,6 +991,35 @@ class Engine(EngineScoreMixin, EngineBase):
return processes, names return processes, names
@staticmethod
def _set_startup_time(
tokenizer_manager: Union[TokenizerManager, MultiTokenizerRouter],
scheduler_init_result: SchedulerInitResult,
startup_tic: float,
) -> None:
startup_time = build_engine_startup_time(
(
info.get("startup_time")
for info in scheduler_init_result.scheduler_infos
),
tokenizer_e2e=time.perf_counter() - startup_tic,
)
tokenizer_manager.set_startup_time(startup_time)
cuda_graph_timings = ", ".join(
f"{phase}={duration:.2f}"
for phase, duration in startup_time["cuda_graph"].items()
)
logger.info(
"Engine startup timings (s): load_weight=%.2f, "
"kv_cache_allocation=%.2f, scheduler_e2e=%.2f, "
"cuda_graph={%s}, tokenizer_e2e=%.2f",
startup_time["load_weight"],
startup_time["kv_cache_allocation"],
startup_time["scheduler_e2e"],
cuda_graph_timings,
startup_time["tokenizer_e2e"],
)
@classmethod @classmethod
def _launch_subprocesses( def _launch_subprocesses(
cls, cls,
@@ -1010,6 +1040,8 @@ class Engine(EngineScoreMixin, EngineBase):
Returns: Returns:
Tuple of (tokenizer_manager, template_manager, port_args, scheduler_init_result, subprocess_watchdog, weight_cache_daemon_procs). Tuple of (tokenizer_manager, template_manager, port_args, scheduler_init_result, subprocess_watchdog, weight_cache_daemon_procs).
""" """
startup_tic = time.perf_counter()
# Configure global environment # Configure global environment
configure_logger(server_args) configure_logger(server_args)
_set_envs_and_config(server_args) _set_envs_and_config(server_args)
@@ -1144,6 +1176,8 @@ class Engine(EngineScoreMixin, EngineBase):
# Wait for the model to finish loading # Wait for the model to finish loading
scheduler_init_result.wait_for_ready() scheduler_init_result.wait_for_ready()
cls._set_startup_time(tokenizer_manager, scheduler_init_result, startup_tic)
# Get back some info from scheduler to tokenizer_manager # Get back some info from scheduler to tokenizer_manager
tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[0][ tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[0][
"max_req_input_len" "max_req_input_len"
@@ -1275,6 +1309,7 @@ class Engine(EngineScoreMixin, EngineBase):
dataclasses.asdict(self.tokenizer_manager.server_args) dataclasses.asdict(self.tokenizer_manager.server_args)
), ),
**self._scheduler_init_result.scheduler_infos[0], **self._scheduler_init_result.scheduler_infos[0],
"startup_time": self.tokenizer_manager.startup_time,
"internal_states": internal_states, "internal_states": internal_states,
"version": __version__, "version": __version__,
} }
+8 -1
View File
@@ -253,6 +253,7 @@ async def init_multi_tokenizer() -> ServerArgs:
) )
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"] tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
tokenizer_manager.set_startup_time(scheduler_info["startup_time"])
set_global_state( set_global_state(
_GlobalState( _GlobalState(
@@ -794,6 +795,7 @@ async def server_info():
dataclasses.asdict(server_args) dataclasses.asdict(server_args)
), ),
**_global_state.scheduler_info, **_global_state.scheduler_info,
"startup_time": _global_state.tokenizer_manager.startup_time,
"internal_states": internal_states, "internal_states": internal_states,
"version": __version__, "version": __version__,
# Structured KV-event publisher descriptor for KV-aware routers. # Structured KV-event publisher descriptor for KV-aware routers.
@@ -2521,7 +2523,12 @@ def _setup_and_run_http_server(
# for other worker processes to read. # for other worker processes to read.
app.is_single_tokenizer_mode = False app.is_single_tokenizer_mode = False
multi_tokenizer_args_shm = write_data_for_multi_tokenizer( multi_tokenizer_args_shm = write_data_for_multi_tokenizer(
port_args, server_args, scheduler_infos[0] port_args,
server_args,
{
**scheduler_infos[0],
"startup_time": tokenizer_manager.startup_time,
},
) )
try: try:
@@ -294,7 +294,6 @@ class MlxModelRunnerStub(ModelRunner):
# No CUDA graphs, no attention backend # No CUDA graphs, no attention backend
self.decode_cuda_graph_runner = None self.decode_cuda_graph_runner = None
self.graph_mem_usage = 0
self.attn_backend = None self.attn_backend = None
self.init_ngram_embedding_manager() self.init_ngram_embedding_manager()
@@ -47,6 +47,7 @@ from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import run_scheduler_process from sglang.srt.managers.scheduler import run_scheduler_process
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
from sglang.srt.observability.startup_time import aggregate_scheduler_startup_times
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.runtime_context import get_exec, publish from sglang.srt.runtime_context import get_exec, publish
from sglang.srt.server_args import ( from sglang.srt.server_args import (
@@ -731,6 +732,9 @@ class DataParallelController:
self.max_total_num_tokens = scheduler_info[0]["max_total_num_tokens"] self.max_total_num_tokens = scheduler_info[0]["max_total_num_tokens"]
self.max_req_input_len = scheduler_info[0]["max_req_input_len"] self.max_req_input_len = scheduler_info[0]["max_req_input_len"]
self.startup_time = aggregate_scheduler_startup_times(
info.get("startup_time") for info in scheduler_info
)
def maybe_external_dp_rank_routing(self, req: Req): def maybe_external_dp_rank_routing(self, req: Req):
if req.routed_dp_rank is not None: if req.routed_dp_rank is not None:
@@ -844,6 +848,7 @@ def run_data_parallel_controller_process(
"status": "ready", "status": "ready",
"max_total_num_tokens": controller.max_total_num_tokens, "max_total_num_tokens": controller.max_total_num_tokens,
"max_req_input_len": controller.max_req_input_len, "max_req_input_len": controller.max_req_input_len,
"startup_time": controller.startup_time,
SCHEDULER_PIDS_ARG: scheduler_pids, SCHEDULER_PIDS_ARG: scheduler_pids,
} }
) )
@@ -440,6 +440,7 @@ class MultiTokenizerRouter:
port_args: PortArgs, port_args: PortArgs,
): ):
self.server_args = server_args self.server_args = server_args
self.startup_time: Optional[Dict[str, Any]] = None
context = zmq.asyncio.Context(3) context = zmq.asyncio.Context(3)
self.recv_from_detokenizer = get_zmq_socket( self.recv_from_detokenizer = get_zmq_socket(
context, zmq.PULL, port_args.tokenizer_ipc_name, True context, zmq.PULL, port_args.tokenizer_ipc_name, True
@@ -479,6 +480,9 @@ class MultiTokenizerRouter:
# Shared socket mapping (both coroutines run on self._loop, so safe) # Shared socket mapping (both coroutines run on self._loop, so safe)
self.socket_mapping = SocketMapping() self.socket_mapping = SocketMapping()
def set_startup_time(self, startup_time: Dict[str, Any]) -> None:
self.startup_time = startup_time
def _run_loop(self): def _run_loop(self):
self._loop.run_forever() self._loop.run_forever()
@@ -792,7 +796,9 @@ def read_from_shared_memory(name: str) -> Any:
def write_data_for_multi_tokenizer( def write_data_for_multi_tokenizer(
port_args: PortArgs, server_args: ServerArgs, scheduler_info: Dict port_args: PortArgs,
server_args: ServerArgs,
scheduler_info: Dict,
): ):
"""Write args information to share memory for multi-tokenizer""" """Write args information to share memory for multi-tokenizer"""
# get main process ID # get main process ID
+92 -33
View File
@@ -218,6 +218,10 @@ from sglang.srt.managers.scheduler_components.load_inquirer import SchedulerLoad
from sglang.srt.managers.scheduler_components.logprob_result_processor import ( from sglang.srt.managers.scheduler_components.logprob_result_processor import (
SchedulerLogprobResultProcessor, SchedulerLogprobResultProcessor,
) )
from sglang.srt.managers.scheduler_components.memory_usage import (
build_memory_usage,
combine_graph_memory_usage,
)
from sglang.srt.managers.scheduler_components.metrics_reporter import ( from sglang.srt.managers.scheduler_components.metrics_reporter import (
RECORD_STEP_TIME, RECORD_STEP_TIME,
PrefillStats, PrefillStats,
@@ -262,6 +266,7 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_time_batch, set_schedule_time_batch,
set_time_batch, set_time_batch,
) )
from sglang.srt.observability.startup_time import build_scheduler_startup_time
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
@@ -378,6 +383,11 @@ class Scheduler(
moe_dp_rank: int, moe_dp_rank: int,
dp_rank: Optional[int], dp_rank: Optional[int],
): ):
# NOTE: KEEP THE FOLLOWING CODE STYLE for this function:
# Keep __init__ as an orchestrator: sequence init_* and maybe_init_* calls
# with minimal glue. Move substantial component-specific logic into
# dedicated methods instead of adding inline blocks here.
self.init_startup_timing_begin()
self.is_initializing = True self.is_initializing = True
# init_soft_watchdog starts a daemon thread that reads these on its first tick. # init_soft_watchdog starts a daemon thread that reads these on its first tick.
self.forward_ct: int = 0 self.forward_ct: int = 0
@@ -521,22 +531,8 @@ class Scheduler(
self.token_to_kv_pool_allocator = result.token_to_kv_pool_allocator self.token_to_kv_pool_allocator = result.token_to_kv_pool_allocator
self.disable_radix_cache = result.disable_radix_cache self.disable_radix_cache = result.disable_radix_cache
self.tree_cache = result.tree_cache self.tree_cache = result.tree_cache
self.emit_metrics_constants()
if _is_npu and is_deepseek_v4( self.maybe_init_hccl_dp_prewarm()
self.tp_worker.model_runner.model_config.hf_config
):
rank = (
self.ps.dp_rank
if self.ps.dp_rank is not None
else self.tp_group.rank_in_group
)
logger.info("HCCL DP prewarm start: rank=%s", rank)
_prewarm_hccl_group(
device=self.tp_group.device,
group=self.tp_group.device_group,
device_module=self.tp_group.device_module,
)
logger.info("HCCL DP prewarm done: rank=%s", rank)
if (c := self.tp_worker.model_runner.canary_manager) is not None: if (c := self.tp_worker.model_runner.canary_manager) is not None:
c.attach_radix_cache(self.tree_cache) c.attach_radix_cache(self.tree_cache)
@@ -645,6 +641,46 @@ class Scheduler(
self.init_batch_result_processor() self.init_batch_result_processor()
self.is_initializing = False self.is_initializing = False
self.init_startup_timing_summary()
def init_startup_timing_begin(self) -> None:
self.scheduler_startup_begin = time.perf_counter()
def init_startup_timing_summary(self) -> None:
self.startup_time = build_scheduler_startup_time(
target_load_weight=self.tp_worker.weight_load_time,
draft_load_weight=(
0.0 if self.draft_worker is None else self.draft_worker.weight_load_time
),
kv_cache_allocation=self.kv_cache_allocation_time,
scheduler_e2e=time.perf_counter() - self.scheduler_startup_begin,
target_cuda_graph=self.tp_worker.graph_time_usage,
draft_cuda_graph=(
None
if self.draft_worker is None
else self.draft_worker.graph_time_usage
),
)
def maybe_init_hccl_dp_prewarm(self) -> None:
if not (
_is_npu
and is_deepseek_v4(self.tp_worker.model_runner.model_config.hf_config)
):
return
rank = (
self.ps.dp_rank
if self.ps.dp_rank is not None
else self.tp_group.rank_in_group
)
logger.info("HCCL DP prewarm start: rank=%s", rank)
_prewarm_hccl_group(
device=self.tp_group.device,
group=self.tp_group.device_group,
device_module=self.tp_group.device_module,
)
logger.info("HCCL DP prewarm done: rank=%s", rank)
def init_zbal_on_npu(self): def init_zbal_on_npu(self):
if _is_npu: if _is_npu:
@@ -941,14 +977,18 @@ class Scheduler(
self.maybe_init_draft_worker() self.maybe_init_draft_worker()
# Prepare KV cache pools for all workers # Prepare KV cache pools for all workers
tic = time.perf_counter()
self.init_memory_pools() self.init_memory_pools()
self.kv_cache_allocation_time = time.perf_counter() - tic
self.init_all_attention_backends() self.init_all_attention_backends()
self.init_all_cuda_graphs() self.init_all_cuda_graphs()
model_runner = self.tp_worker.model_runner model_runner = self.tp_worker.model_runner
if model_runner.token_to_kv_pool.post_capture_active: if model_runner.token_to_kv_pool.post_capture_active:
tic = time.perf_counter()
model_runner.post_capture_resize_kv_pool() model_runner.post_capture_resize_kv_pool()
self.kv_cache_allocation_time += time.perf_counter() - tic
if ( if (
get_exec().moe.elastic_ep_backend is not None get_exec().moe.elastic_ep_backend is not None
@@ -1021,7 +1061,7 @@ class Scheduler(
set_random_seed(self.random_seed) set_random_seed(self.random_seed)
# Print debug info # Print debug info
avail_mem = get_available_gpu_memory( self.startup_available_gpu_memory_gb = get_available_gpu_memory(
self.device, self.ps.gpu_id, empty_cache=False self.device, self.ps.gpu_id, empty_cache=False
) )
if self.ps.tp_rank == 0: if self.ps.tp_rank == 0:
@@ -1031,22 +1071,35 @@ class Scheduler(
f"max_prefill_tokens={self.max_prefill_tokens}, " f"max_prefill_tokens={self.max_prefill_tokens}, "
f"max_running_requests={self.max_running_requests}, " f"max_running_requests={self.max_running_requests}, "
f"context_len={self.model_config.context_len}, " f"context_len={self.model_config.context_len}, "
f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB" f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}="
f"{self.startup_available_gpu_memory_gb:.2f} GB"
) )
if get_observability().enable_metrics: def emit_metrics_constants(self) -> None:
if not get_observability().enable_metrics:
return
self.metrics_collector.emit_constants( self.metrics_collector.emit_constants(
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
# TODO: max_running_requests_under_SLO has no setter — dead chain. max_total_num_tokens_swa=self.swa_tokens_per_layer,
max_running_requests_under_SLO=getattr( weight_memory_usage_gb=self.tp_worker.model_runner.weight_load_mem_usage,
self, "max_running_requests_under_SLO", None kv_cache_memory_usage_gb=(
self.token_to_kv_pool_allocator.get_kvcache().mem_usage
), ),
engine_startup_time=0.0, graph_memory_usage_gb=combine_graph_memory_usage(
engine_load_weights_time=0.0, self.tp_worker.graph_memory_usage,
(
None
if self.draft_worker is None
else self.draft_worker.graph_memory_usage
),
),
# TODO: max_running_requests_under_SLO has no setter — dead chain.
max_running_requests_under_SLO=None,
page_size=self.page_size, page_size=self.page_size,
num_pages=self.max_total_num_tokens // self.page_size, num_pages=self.max_total_num_tokens // self.page_size,
context_len=self.model_config.context_len, context_len=self.model_config.context_len,
startup_available_gpu_memory_gb=avail_mem, startup_available_gpu_memory_gb=self.startup_available_gpu_memory_gb,
) )
def init_hisparse_coordinator(self) -> None: def init_hisparse_coordinator(self) -> None:
@@ -1561,6 +1614,7 @@ class Scheduler(
"status": "ready", "status": "ready",
"max_total_num_tokens": self.max_total_num_tokens, "max_total_num_tokens": self.max_total_num_tokens,
"max_req_input_len": self.max_req_input_len, "max_req_input_len": self.max_req_input_len,
"startup_time": self.startup_time,
} }
return result_dict return result_dict
@@ -4113,14 +4167,19 @@ class Scheduler(
# readback reflects values changed via /set_internal_state, not startup. # readback reflects values changed via /set_internal_state, not startup.
ret = get_context().resolved_server_args_dict() ret = get_context().resolved_server_args_dict()
ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput
ret["memory_usage"] = { draft_graph_memory_usage = (
"weight": round(self.tp_worker.model_runner.weight_load_mem_usage, 2), None if self.draft_worker is None else self.draft_worker.graph_memory_usage
"kvcache": round( )
self.token_to_kv_pool_allocator.get_kvcache().mem_usage, 2 ret["memory_usage"] = build_memory_usage(
), weight_gb=self.tp_worker.model_runner.weight_load_mem_usage,
"token_capacity": int(self.max_total_num_tokens), kv_cache_gb=self.token_to_kv_pool_allocator.get_kvcache().mem_usage,
"graph": round(self.tp_worker.model_runner.graph_mem_usage, 2), startup_available_gb=self.startup_available_gpu_memory_gb,
} token_capacity=self.max_total_num_tokens,
token_capacity_swa=self.swa_tokens_per_layer,
target_graph_memory_usage=self.tp_worker.graph_memory_usage,
draft_graph_memory_usage=draft_graph_memory_usage,
)
ret["startup_time"] = self.startup_time
ret["effective_max_running_requests_per_dp"] = self.max_running_requests ret["effective_max_running_requests_per_dp"] = self.max_running_requests
if get_exec().moe.elastic_ep_backend is not None: if get_exec().moe.elastic_ep_backend is not None:
@@ -136,7 +136,7 @@ class SchedulerLoadInquirer:
kv_cache_gb=round( kv_cache_gb=round(
self.token_to_kv_pool_allocator.get_kvcache().mem_usage, 3 self.token_to_kv_pool_allocator.get_kvcache().mem_usage, 3
), ),
graph_gb=round(self.tp_worker.model_runner.graph_mem_usage, 3), graph_gb=round(sum(self.tp_worker.graph_memory_usage.values()), 3),
token_capacity=int(self.max_total_num_tokens), token_capacity=int(self.max_total_num_tokens),
) )
except (AttributeError, TypeError) as e: except (AttributeError, TypeError) as e:
@@ -0,0 +1,41 @@
from __future__ import annotations
from collections.abc import Mapping
from sglang.srt.model_executor.graph_memory_usage import merge_graph_memory_usage
def combine_graph_memory_usage(
target: Mapping[str, float] | None,
draft: Mapping[str, float] | None,
) -> dict[str, float]:
return merge_graph_memory_usage(target, draft)
def build_memory_usage(
*,
weight_gb: float,
kv_cache_gb: float,
startup_available_gb: float,
token_capacity: int,
token_capacity_swa: int | None,
target_graph_memory_usage: Mapping[str, float] | None,
draft_graph_memory_usage: Mapping[str, float] | None,
) -> dict:
graph_memory_usage = combine_graph_memory_usage(
target_graph_memory_usage,
draft_graph_memory_usage,
)
return {
"weight": round(weight_gb, 3),
"kvcache": round(kv_cache_gb, 3),
"startup_available": round(startup_available_gb, 3),
"token_capacity": int(token_capacity),
"token_capacity_swa": (
None if token_capacity_swa is None else int(token_capacity_swa)
),
"graph": {
phase: round(memory_gb, 3)
for phase, memory_gb in graph_memory_usage.items()
},
}
@@ -391,6 +391,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
): ):
# Parse args # Parse args
self.server_args = server_args self.server_args = server_args
self.startup_time: Optional[Dict[str, Any]] = None
self._config_updates: List[Tuple[str, Dict[str, Any]]] = [] self._config_updates: List[Tuple[str, Dict[str, Any]]] = []
self.elastic_worker_count = server_args.dp_size self.elastic_worker_count = server_args.dp_size
self.elastic_pending_ep_size = None self.elastic_pending_ep_size = None
@@ -701,6 +702,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(), test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
) )
def set_startup_time(self, startup_time: Dict[str, Any]) -> None:
self.startup_time = startup_time
if self.enable_metrics:
self.metrics_collector.emit_startup_time(startup_time)
def init_request_dispatcher(self): def init_request_dispatcher(self):
self._result_dispatcher = TypeBasedDispatcher( self._result_dispatcher = TypeBasedDispatcher(
[ [
+21
View File
@@ -46,6 +46,10 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.graph_memory_usage import (
merge_graph_memory_usage,
merge_graph_time_usage,
)
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -97,6 +101,23 @@ class BaseTpWorker(ABC):
self.model_runner.swa_max_total_num_tokens, self.model_runner.swa_max_total_num_tokens,
) )
@property
def graph_memory_usage(self) -> dict[str, float]:
runners = self.model_runner_list or [self.model_runner]
return merge_graph_memory_usage(
*(runner.graph_memory_usage for runner in runners)
)
@property
def graph_time_usage(self) -> dict[str, float]:
runners = self.model_runner_list or [self.model_runner]
return merge_graph_time_usage(*(runner.graph_time_usage for runner in runners))
@property
def weight_load_time(self) -> float:
runners = self.model_runner_list or [self.model_runner]
return sum(runner.weight_load_time for runner in runners)
def get_pad_input_ids_func(self): def get_pad_input_ids_func(self):
return getattr(self.model_runner.model, "pad_input_ids", None) return getattr(self.model_runner.model, "pad_input_ids", None)
+14 -2
View File
@@ -1593,6 +1593,7 @@ class KVCache(abc.ABC):
enable_memory_saver: bool, enable_memory_saver: bool,
start_layer: Optional[int] = None, start_layer: Optional[int] = None,
end_layer: Optional[int] = None, end_layer: Optional[int] = None,
allocation_label: Optional[str] = None,
): ):
self.size = size self.size = size
self.page_size = page_size self.page_size = page_size
@@ -1606,6 +1607,7 @@ class KVCache(abc.ABC):
self.layer_num = layer_num self.layer_num = layer_num
self.start_layer = start_layer or 0 self.start_layer = start_layer or 0
self.end_layer = end_layer or layer_num - 1 self.end_layer = end_layer or layer_num - 1
self.allocation_label = allocation_label
self.memory_saver_adapter = TorchMemorySaverAdapter.create( self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver enable=enable_memory_saver
) )
@@ -1626,19 +1628,27 @@ class KVCache(abc.ABC):
"""Common logging and mem_usage computation for KV cache allocation. """Common logging and mem_usage computation for KV cache allocation.
Supports both tuple (K, V) size returns and single KV size returns. Supports both tuple (K, V) size returns and single KV size returns.
""" """
cache_name = (
f"{self.allocation_label} KV Cache"
if self.allocation_label is not None
else "KV Cache"
)
kv_size_bytes = self.get_kv_size_bytes() kv_size_bytes = self.get_kv_size_bytes()
if isinstance(kv_size_bytes, tuple): if isinstance(kv_size_bytes, tuple):
k_size, v_size = kv_size_bytes k_size, v_size = kv_size_bytes
k_size_GB = k_size / GB k_size_GB = k_size / GB
v_size_GB = v_size / GB v_size_GB = v_size / GB
logger.info( logger.info(
f"KV Cache is allocated. dtype: {self.dtype}, #tokens: {num_tokens}, K size: {k_size_GB:.2f} GB, V size: {v_size_GB:.2f} GB" f"{cache_name} is allocated. dtype: {self.dtype}, "
f"#tokens: {num_tokens}, K size: {k_size_GB:.2f} GB, "
f"V size: {v_size_GB:.2f} GB"
) )
self.mem_usage = k_size_GB + v_size_GB self.mem_usage = k_size_GB + v_size_GB
else: else:
kv_size_GB = kv_size_bytes / GB kv_size_GB = kv_size_bytes / GB
logger.info( logger.info(
f"KV Cache is allocated. dtype: {self.dtype}, #tokens: {num_tokens}, KV size: {kv_size_GB:.2f} GB" f"{cache_name} is allocated. dtype: {self.dtype}, "
f"#tokens: {num_tokens}, KV size: {kv_size_GB:.2f} GB"
) )
self.mem_usage = kv_size_GB self.mem_usage = kv_size_GB
@@ -1721,6 +1731,7 @@ class MHATokenToKVPool(KVCache):
kv_cache_layout: Optional[str] = None, kv_cache_layout: Optional[str] = None,
quant_method=None, quant_method=None,
post_capture_active: bool = False, post_capture_active: bool = False,
allocation_label: Optional[str] = None,
): ):
self.k_buffer = None self.k_buffer = None
self.v_buffer = None self.v_buffer = None
@@ -1737,6 +1748,7 @@ class MHATokenToKVPool(KVCache):
enable_memory_saver, enable_memory_saver,
start_layer, start_layer,
end_layer, end_layer,
allocation_label,
) )
self.post_capture_active = post_capture_active self.post_capture_active = post_capture_active
self._post_capture_owner = None self._post_capture_owner = None
+12 -9
View File
@@ -57,19 +57,22 @@ class SWAKVPool(BaseSWAKVPool):
maybe_init_custom_mem_pool(device=self.device) maybe_init_custom_mem_pool(device=self.device)
) )
self.swa_kv_pool = token_to_kv_pool_class( full_pool_kwargs = kwargs.copy()
size=size_swa, full_pool_kwargs.pop("swa_head_num", None)
dtype=dtype, full_pool_kwargs.pop("swa_head_dim", None)
layer_num=self.swa_layer_nums, full_pool_kwargs.pop("swa_v_head_dim", None)
**kwargs,
)
kwargs.pop("swa_head_num", None)
kwargs.pop("swa_head_dim", None)
kwargs.pop("swa_v_head_dim", None)
self.full_kv_pool = token_to_kv_pool_class( self.full_kv_pool = token_to_kv_pool_class(
size=size, size=size,
dtype=dtype, dtype=dtype,
layer_num=self.full_layer_nums, layer_num=self.full_layer_nums,
allocation_label="Full",
**full_pool_kwargs,
)
self.swa_kv_pool = token_to_kv_pool_class(
size=size_swa,
dtype=dtype,
layer_num=self.swa_layer_nums,
allocation_label="SWA",
**kwargs, **kwargs,
) )
# {layer_id: (index, is_swa_layer)} # {layer_id: (index, is_swa_layer)}
@@ -1265,11 +1265,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
dp_padding_mode = DpPaddingMode.SUM_LEN dp_padding_mode = DpPaddingMode.SUM_LEN
# Prefill breakable CUDA graph requires every DP rank to run the SAME # Prefill breakable CUDA graph requires every DP rank to run the SAME
# captured shape. Under SUM_LEN each rank pads to its own local token # captured shape. Under SUM_LEN each rank pads to its own local token
# count and can select a different capture bucket, so the in-graph DP # count and can select a different capture bucket. This mismatches the
# collectives (all_gather / reduce_scatter) mismatch across ranks and # rank-coupled communication geometry: DP gather/combine uses
# corrupt the output. Force MAX_LEN so every rank pads to the global # all_gather_into_tensor / reduce_scatter_tensor, while MoE backends may
# max and picks the same bucket (mirrors the decode cuda graph # use A2A dispatch/combine. Force MAX_LEN so every rank pads to the global
# contract, which always runs MAX_LEN). # max and picks the same bucket.
# #
# Only force MAX_LEN when the batch fits a captured breakable prefill # Only force MAX_LEN when the batch fits a captured breakable prefill
# graph; larger prefills fall back to eager and keep the # graph; larger prefills fall back to eager and keep the
@@ -0,0 +1,66 @@
from __future__ import annotations
from collections.abc import Mapping
GRAPH_MEMORY_USAGE_KEYS = (
"prefill",
"decode",
"target_verify",
"draft_prefill",
"draft_decode",
"draft_extend",
)
def empty_graph_memory_usage() -> dict[str, float]:
return dict.fromkeys(GRAPH_MEMORY_USAGE_KEYS, 0.0)
def merge_graph_memory_usage(
*usages: Mapping[str, float] | None,
) -> dict[str, float]:
"""Sum graph-capture memory by phase and keep a stable base schema."""
merged = empty_graph_memory_usage()
for usage in usages:
if usage is None:
continue
for phase, value in usage.items():
merged[phase] = merged.get(phase, 0.0) + value
return merged
def replace_graph_memory_usage(
current: Mapping[str, float] | None,
replacement: Mapping[str, float],
*,
phases: tuple[str, ...],
) -> dict[str, float]:
"""Replace one capture family while preserving measurements for the rest."""
updated = merge_graph_memory_usage(current)
for phase in phases:
updated[phase] = 0.0
updated.update(replacement)
return updated
def merge_graph_time_usage(
*usages: Mapping[str, float] | None,
) -> dict[str, float]:
return merge_graph_memory_usage(*usages)
def empty_graph_time_usage() -> dict[str, float]:
return empty_graph_memory_usage()
def replace_graph_time_usage(
current: Mapping[str, float] | None,
replacement: Mapping[str, float],
*,
phases: tuple[str, ...],
) -> dict[str, float]:
return replace_graph_memory_usage(
current,
replacement,
phases=phases,
)
@@ -103,6 +103,10 @@ from sglang.srt.model_executor.forward_context import (
forward_context, forward_context,
has_forward_context, has_forward_context,
) )
from sglang.srt.model_executor.graph_memory_usage import (
replace_graph_memory_usage,
replace_graph_time_usage,
)
from sglang.srt.model_executor.model_runner_components import misc_utils from sglang.srt.model_executor.model_runner_components import misc_utils
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import ( from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
build_attention_backends, build_attention_backends,
@@ -310,6 +314,8 @@ class ModelRunner:
self.draft_model_idx = draft_model_idx self.draft_model_idx = draft_model_idx
self.enable_hisparse = server_args.enable_hisparse self.enable_hisparse = server_args.enable_hisparse
self.init_startup_observability()
self.init_remote_instance_weight_transporter() self.init_remote_instance_weight_transporter()
self.init_msprobe() self.init_msprobe()
@@ -411,6 +417,11 @@ class ModelRunner:
self.init_weight_updater() self.init_weight_updater()
self.init_weight_exporter() self.init_weight_exporter()
def init_startup_observability(self) -> None:
self.weight_load_time = 0.0
self.graph_memory_usage: dict[str, float] = {}
self.graph_time_usage: dict[str, float] = {}
def _initialize_elastic_ep_joiner(self) -> None: def _initialize_elastic_ep_joiner(self) -> None:
if not ( if not (
get_exec().moe.elastic_ep_backend is not None get_exec().moe.elastic_ep_backend is not None
@@ -925,9 +936,10 @@ class ModelRunner:
model_runner=self, capture_decode_cuda_graph=capture_decode_cuda_graph model_runner=self, capture_decode_cuda_graph=capture_decode_cuda_graph
) )
self.eager_runner = capture.eager_runner self.eager_runner = capture.eager_runner
self.prefill_cuda_graph_runner = capture.prefill_runner self.prefill_cuda_graph_runner = capture.prefill.runner
self.decode_cuda_graph_runner = capture.decode.runner self.decode_cuda_graph_runner = capture.decode.runner
self.graph_mem_usage = capture.decode.graph_mem_usage self.graph_memory_usage = capture.memory_usage
self.graph_time_usage = capture.time_usage
def init_routed_experts_capturer(self): def init_routed_experts_capturer(self):
if self.is_draft_worker: if self.is_draft_worker:
@@ -1068,13 +1080,14 @@ class ModelRunner:
after_avail_memory = get_available_gpu_memory(self.device, self.gpu_id) after_avail_memory = get_available_gpu_memory(self.device, self.gpu_id)
self.weight_load_mem_usage = before_avail_memory - after_avail_memory self.weight_load_mem_usage = before_avail_memory - after_avail_memory
self.weight_load_time = time.perf_counter() - tic_total
# Get quantization config from ModelConfig # Get quantization config from ModelConfig
# This handles both config.json (standard) and hf_quant_config.json (ModelOpt) # This handles both config.json (standard) and hf_quant_config.json (ModelOpt)
quant_str = self.model_config.get_quantization_config_log_str() quant_str = self.model_config.get_quantization_config_log_str()
logger.info( logger.info(
f"Load weight end. " f"Load weight end. "
f"elapsed={time.perf_counter() - tic_total:.2f} s, " f"elapsed={self.weight_load_time:.2f} s, "
f"type={type(self.model).__name__}, " f"type={type(self.model).__name__}, "
f"{quant_str + ', ' if quant_str else ''}" f"{quant_str + ', ' if quant_str else ''}"
f"avail mem={after_avail_memory:.2f} GB, " f"avail mem={after_avail_memory:.2f} GB, "
@@ -1212,18 +1225,37 @@ class ModelRunner:
def init_decode_cuda_graph(self): def init_decode_cuda_graph(self):
self.decode_cuda_graph_runner = None self.decode_cuda_graph_runner = None
self.graph_mem_usage = 0
capture = capture_decode_graph(model_runner=self) capture = capture_decode_graph(model_runner=self)
self.decode_cuda_graph_runner = capture.runner self.decode_cuda_graph_runner = capture.runner
self.graph_mem_usage = capture.graph_mem_usage self.graph_memory_usage = replace_graph_memory_usage(
self.graph_memory_usage,
capture.memory_usage,
phases=("decode", "target_verify", "draft_decode"),
)
self.graph_time_usage = replace_graph_time_usage(
self.graph_time_usage,
capture.time_usage,
phases=("decode", "target_verify", "draft_decode"),
)
def init_prefill_cuda_graph(self, force_for_draft_worker: bool = False): def init_prefill_cuda_graph(self, force_for_draft_worker: bool = False):
self.prefill_cuda_graph_runner = None self.prefill_cuda_graph_runner = None
self.prefill_cuda_graph_runner = capture_prefill_graph( capture = capture_prefill_graph(
model_runner=self, model_runner=self,
eager_runner=self.eager_runner, eager_runner=self.eager_runner,
force_for_draft_worker=force_for_draft_worker, force_for_draft_worker=force_for_draft_worker,
) )
self.prefill_cuda_graph_runner = capture.runner
self.graph_memory_usage = replace_graph_memory_usage(
self.graph_memory_usage,
capture.memory_usage,
phases=("prefill", "draft_prefill"),
)
self.graph_time_usage = replace_graph_time_usage(
self.graph_time_usage,
capture.time_usage,
phases=("prefill", "draft_prefill"),
)
def init_threads_binding(self): def init_threads_binding(self):
self.local_omp_cpuid = numa_utils.init_threads_binding( self.local_omp_cpuid = numa_utils.init_threads_binding(
@@ -23,6 +23,10 @@ from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
get_server_return_hidden_states_mode, get_server_return_hidden_states_mode,
) )
from sglang.srt.model_executor.graph_memory_usage import (
merge_graph_memory_usage,
merge_graph_time_usage,
)
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_components.layer_setup import ( from sglang.srt.model_executor.model_runner_components.layer_setup import (
@@ -63,15 +67,39 @@ def should_skip_auto_prefill_cuda_graph_for_memory(
) )
class DecodeGraphCapture(msgspec.Struct, frozen=True, kw_only=True): class GraphCapture(msgspec.Struct, frozen=True, kw_only=True):
runner: Optional[BaseRunner] runner: Optional[BaseRunner]
graph_mem_usage: float memory_phase: str
memory_usage_gb: float
capture_time: float
@property
def memory_usage(self) -> dict[str, float]:
return {self.memory_phase: self.memory_usage_gb}
@property
def time_usage(self) -> dict[str, float]:
return {self.memory_phase: self.capture_time}
class CudaGraphsCapture(msgspec.Struct, frozen=True, kw_only=True): class CudaGraphsCapture(msgspec.Struct, frozen=True, kw_only=True):
eager_runner: EagerRunner eager_runner: EagerRunner
prefill_runner: Optional[BaseRunner] prefill: GraphCapture
decode: DecodeGraphCapture decode: GraphCapture
@property
def memory_usage(self) -> dict[str, float]:
return merge_graph_memory_usage(
self.prefill.memory_usage,
self.decode.memory_usage,
)
@property
def time_usage(self) -> dict[str, float]:
return merge_graph_time_usage(
self.prefill.time_usage,
self.decode.time_usage,
)
def capture_cuda_graphs( def capture_cuda_graphs(
@@ -100,11 +128,17 @@ def capture_cuda_graphs(
# cuda-graph capture: prefill before decode, so both coalesce onto the # cuda-graph capture: prefill before decode, so both coalesce onto the
# eager buffer allocated above. (capture_prefill_graph routes prefill # eager buffer allocated above. (capture_prefill_graph routes prefill
# to the eager runner when the prefill graph is disabled.) # to the eager runner when the prefill graph is disabled.)
prefill_runner = capture_prefill_graph( prefill = capture_prefill_graph(
model_runner=model_runner, eager_runner=eager_runner model_runner=model_runner, eager_runner=eager_runner
) )
decode = DecodeGraphCapture(runner=None, graph_mem_usage=0) decode_phase = "draft_decode" if model_runner.is_draft_worker else "decode"
decode = GraphCapture(
runner=None,
memory_phase=decode_phase,
memory_usage_gb=0,
capture_time=0,
)
if capture_decode_cuda_graph: if capture_decode_cuda_graph:
if model_runner.device in ("cuda", "musa", "cpu", "npu", "xpu"): if model_runner.device in ("cuda", "musa", "cpu", "npu", "xpu"):
decode = capture_decode_graph(model_runner=model_runner) decode = capture_decode_graph(model_runner=model_runner)
@@ -113,7 +147,12 @@ def capture_cuda_graphs(
): ):
decode = capture_decode_graph(model_runner=model_runner) decode = capture_decode_graph(model_runner=model_runner)
else: else:
decode = DecodeGraphCapture(runner=eager_runner, graph_mem_usage=0) decode = GraphCapture(
runner=eager_runner,
memory_phase=decode_phase,
memory_usage_gb=0,
capture_time=0,
)
# Register forward hooks AFTER cuda-graph capture so their tensor ops are # Register forward hooks AFTER cuda-graph capture so their tensor ops are
# not traced into any captured graph — capture stays hook-free and hooks # not traced into any captured graph — capture stays hook-free and hooks
@@ -134,9 +173,7 @@ def capture_cuda_graphs(
if model_runner.canary_manager is not None and not model_runner.is_draft_worker: if model_runner.canary_manager is not None and not model_runner.is_draft_worker:
model_runner.canary_manager.mark_init_finished() model_runner.canary_manager.mark_init_finished()
return CudaGraphsCapture( return CudaGraphsCapture(eager_runner=eager_runner, prefill=prefill, decode=decode)
eager_runner=eager_runner, prefill_runner=prefill_runner, decode=decode
)
def capture_prefill_graph( def capture_prefill_graph(
@@ -144,8 +181,22 @@ def capture_prefill_graph(
model_runner: ModelRunner, model_runner: ModelRunner,
eager_runner: EagerRunner, eager_runner: EagerRunner,
force_for_draft_worker: bool = False, force_for_draft_worker: bool = False,
) -> Optional[BaseRunner]: ) -> GraphCapture:
"""Initialize prefill CUDA graph runner.""" """Initialize a prefill graph and return its startup resource usage."""
memory_phase = "draft_prefill" if model_runner.is_draft_worker else "prefill"
def result(
runner: Optional[BaseRunner],
memory_usage_gb: float = 0,
capture_time: float = 0,
) -> GraphCapture:
return GraphCapture(
runner=runner,
memory_phase=memory_phase,
memory_usage_gb=memory_usage_gb,
capture_time=capture_time,
)
if check_cuda_graph_backend(Phase.PREFILL, Backend.DISABLED): if check_cuda_graph_backend(Phase.PREFILL, Backend.DISABLED):
logger.info( logger.info(
@@ -157,14 +208,14 @@ def capture_prefill_graph(
# EagerRunner (its can_run_graph returns False, so _forward_raw's # EagerRunner (its can_run_graph returns False, so _forward_raw's
# extend branch falls through to the eager path). # extend branch falls through to the eager path).
if not model_runner.is_draft_worker: if not model_runner.is_draft_worker:
return eager_runner return result(eager_runner)
return None return result(None)
# Draft models skip here during __init__; the eagle worker calls # Draft models skip here during __init__; the eagle worker calls
# this method explicitly (force_for_draft_worker=True) after # this method explicitly (force_for_draft_worker=True) after
# init_lm_head so graphs capture the final embedding weights. # init_lm_head so graphs capture the final embedding weights.
if model_runner.is_draft_worker and not force_for_draft_worker: if model_runner.is_draft_worker and not force_for_draft_worker:
return None return result(None)
# Skip prefill CG for EAGLE target on tc_piecewise when the fixed server # Skip prefill CG for EAGLE target on tc_piecewise when the fixed server
# capture ceiling is below FULL. EAGLE target prefill requests FULL, so a # capture ceiling is below FULL. EAGLE target prefill requests FULL, so a
@@ -183,7 +234,7 @@ def capture_prefill_graph(
"Disable prefill CUDA graph for EAGLE target on tc_piecewise " "Disable prefill CUDA graph for EAGLE target on tc_piecewise "
"to avoid FP4/MoE decode-replay corruption (#28386)." "to avoid FP4/MoE decode-replay corruption (#28386)."
) )
return eager_runner return result(eager_runner)
if ( if (
model_runner.server_args.enable_lora model_runner.server_args.enable_lora
@@ -194,7 +245,7 @@ def capture_prefill_graph(
"configuration does not support it (unsupported LoRA backend, " "configuration does not support it (unsupported LoRA backend, "
"MoE LoRA, or DP attention)." "MoE LoRA, or DP attention)."
) )
return eager_runner return result(eager_runner)
# Resolve the decoder once. Some VLM wrappers (for example Kimi-VL) # Resolve the decoder once. Some VLM wrappers (for example Kimi-VL)
# expose it as ``language_model`` rather than ``model``. # expose it as ``language_model`` rather than ``model``.
@@ -204,12 +255,12 @@ def capture_prefill_graph(
logger.warning( logger.warning(
"Disable prefill CUDA graph because the model is not a language model" "Disable prefill CUDA graph because the model is not a language model"
) )
return None return result(None)
# Disable prefill CUDA graph for non capture size # Disable prefill CUDA graph for non capture size
if not model_runner.server_args.cuda_graph_config.prefill.bs: if not model_runner.server_args.cuda_graph_config.prefill.bs:
logger.warning("Disable prefill CUDA graph because the capture size is not set") logger.warning("Disable prefill CUDA graph because the capture size is not set")
return None return result(None)
prefill_config = model_runner.server_args.cuda_graph_config.prefill prefill_config = model_runner.server_args.cuda_graph_config.prefill
prefill_backend = prefill_config.backend prefill_backend = prefill_config.backend
@@ -272,7 +323,7 @@ def capture_prefill_graph(
logger.warning( logger.warning(
"Disable prefill CUDA graph because the model does not have a 'layers' attribute" "Disable prefill CUDA graph because the model does not have a 'layers' attribute"
) )
return None return result(None)
( (
model_runner.attention_layers, model_runner.attention_layers,
@@ -288,7 +339,7 @@ def capture_prefill_graph(
logger, logger,
"Disable prefill CUDA graph because some layers do not apply Standard GQA", "Disable prefill CUDA graph because some layers do not apply Standard GQA",
) )
return None return result(None)
tic = time.perf_counter() tic = time.perf_counter()
before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id) before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
@@ -304,7 +355,7 @@ def capture_prefill_graph(
before_mem, before_mem,
_MIN_AUTO_PREFILL_CUDA_GRAPH_FREE_MEMORY_GB, _MIN_AUTO_PREFILL_CUDA_GRAPH_FREE_MEMORY_GB,
) )
return eager_runner return result(eager_runner)
role = "draft" if model_runner.is_draft_worker else "target" role = "draft" if model_runner.is_draft_worker else "target"
capture_name = f"{role} prefill" capture_name = f"{role} prefill"
@@ -318,17 +369,29 @@ def capture_prefill_graph(
after_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id) after_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
mem_usage = before_mem - after_mem mem_usage = before_mem - after_mem
capture_time = time.perf_counter() - tic
logger.info( logger.info(
f"Capture {capture_name} CUDA graph end. " f"Capture {capture_name} CUDA graph end. "
f"elapsed={time.perf_counter() - tic:.2f} s, " f"elapsed={capture_time:.2f} s, "
f"mem usage={mem_usage:.2f} GB, avail mem={after_mem:.2f} GB." f"mem usage={mem_usage:.2f} GB, avail mem={after_mem:.2f} GB."
) )
return prefill_runner return result(prefill_runner, mem_usage, capture_time)
def capture_decode_graph(*, model_runner: ModelRunner) -> DecodeGraphCapture: def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
"""Capture device graphs.""" """Capture device graphs."""
no_capture = DecodeGraphCapture(runner=None, graph_mem_usage=0) if model_runner.is_draft_worker:
memory_phase = "draft_decode"
elif model_runner.spec_algorithm.is_speculative():
memory_phase = "target_verify"
else:
memory_phase = "decode"
no_capture = GraphCapture(
runner=None,
memory_phase=memory_phase,
memory_usage_gb=0,
capture_time=0,
)
if not model_runner.is_generation: if not model_runner.is_generation:
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models # TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
@@ -384,10 +447,16 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> DecodeGraphCapture:
runner = graph_runners[model_runner.device](model_runner) runner = graph_runners[model_runner.device](model_runner)
after_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id) after_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
graph_mem_usage = before_mem - after_mem memory_usage_gb = before_mem - after_mem
capture_time = time.perf_counter() - tic
logger.info( logger.info(
f"Capture {capture_name} {graph_backend[model_runner.device]} end. " f"Capture {capture_name} {graph_backend[model_runner.device]} end. "
f"elapsed={time.perf_counter() - tic:.2f} s, " f"elapsed={capture_time:.2f} s, "
f"mem usage={graph_mem_usage:.2f} GB, avail mem={after_mem:.2f} GB." f"mem usage={memory_usage_gb:.2f} GB, avail mem={after_mem:.2f} GB."
)
return GraphCapture(
runner=runner,
memory_phase=memory_phase,
memory_usage_gb=memory_usage_gb,
capture_time=capture_time,
) )
return DecodeGraphCapture(runner=runner, graph_mem_usage=graph_mem_usage)
@@ -21,7 +21,7 @@ import os
import time import time
from collections import Counter from collections import Counter
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Union from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, Set, Union
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -995,24 +995,36 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labelnames=labels.keys(), labelnames=labels.keys(),
multiprocess_mode="mostrecent", multiprocess_mode="mostrecent",
) )
self.max_total_num_tokens_swa = Gauge(
name="sglang:max_total_num_tokens_swa",
documentation="Maximum total number of tokens in the SWA KV cache pool.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.weight_memory_usage_gb = Gauge(
name="sglang:weight_memory_usage_gb",
documentation="Memory used by model weights in GB.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.kv_cache_memory_usage_gb = Gauge(
name="sglang:kv_cache_memory_usage_gb",
documentation="Memory used by the KV cache pools in GB.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.graph_memory_usage_gb = Gauge(
name="sglang:graph_memory_usage_gb",
documentation="Memory used by captured device graphs in GB.",
labelnames=list(labels.keys()) + ["phase"],
multiprocess_mode="mostrecent",
)
self.max_running_requests_under_SLO = Gauge( self.max_running_requests_under_SLO = Gauge(
name="sglang:max_running_requests_under_SLO", name="sglang:max_running_requests_under_SLO",
documentation="The maximum number of running requests under SLO.", documentation="The maximum number of running requests under SLO.",
labelnames=labels.keys(), labelnames=labels.keys(),
multiprocess_mode="mostrecent", multiprocess_mode="mostrecent",
) )
self.engine_startup_time = Gauge(
name="sglang:engine_startup_time",
documentation="The time taken for the engine to start up.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.engine_load_weights_time = Gauge(
name="sglang:engine_load_weights_time",
documentation="The time taken for the engine to load weights.",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.page_size = Gauge( self.page_size = Gauge(
name="sglang:page_size", name="sglang:page_size",
documentation="KV cache page size in tokens.", documentation="KV cache page size in tokens.",
@@ -1401,21 +1413,30 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
def emit_constants( def emit_constants(
self, self,
max_total_num_tokens: int, max_total_num_tokens: int,
max_total_num_tokens_swa: Optional[int],
weight_memory_usage_gb: float,
kv_cache_memory_usage_gb: float,
graph_memory_usage_gb: Mapping[str, float],
max_running_requests_under_SLO: Optional[int], max_running_requests_under_SLO: Optional[int],
engine_startup_time: float,
engine_load_weights_time: float,
page_size: int, page_size: int,
num_pages: int, num_pages: int,
context_len: int, context_len: int,
startup_available_gpu_memory_gb: float, startup_available_gpu_memory_gb: float,
) -> None: ) -> None:
self._log_gauge(self.max_total_num_tokens, max_total_num_tokens) self._log_gauge(self.max_total_num_tokens, max_total_num_tokens)
if max_total_num_tokens_swa is not None:
self._log_gauge(self.max_total_num_tokens_swa, max_total_num_tokens_swa)
self._log_gauge(self.weight_memory_usage_gb, weight_memory_usage_gb)
self._log_gauge(self.kv_cache_memory_usage_gb, kv_cache_memory_usage_gb)
for phase, memory_usage_gb in graph_memory_usage_gb.items():
self.graph_memory_usage_gb.labels(
**self.labels,
phase=phase,
).set(memory_usage_gb)
if max_running_requests_under_SLO is not None: if max_running_requests_under_SLO is not None:
self._log_gauge( self._log_gauge(
self.max_running_requests_under_SLO, max_running_requests_under_SLO self.max_running_requests_under_SLO, max_running_requests_under_SLO
) )
self._log_gauge(self.engine_startup_time, engine_startup_time)
self._log_gauge(self.engine_load_weights_time, engine_load_weights_time)
self._log_gauge(self.page_size, page_size) self._log_gauge(self.page_size, page_size)
self._log_gauge(self.num_pages, num_pages) self._log_gauge(self.num_pages, num_pages)
self._log_gauge(self.context_len, context_len) self._log_gauge(self.context_len, context_len)
@@ -1435,13 +1456,28 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
) -> None: ) -> None:
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR` # We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
from prometheus_client import Counter as _PromCounter from prometheus_client import Counter as _PromCounter
from prometheus_client import Gauge as _PromGauge
from prometheus_client import Histogram as _PromHistogram from prometheus_client import Histogram as _PromHistogram
Counter = self._counter_cls or _PromCounter Counter = self._counter_cls or _PromCounter
Gauge = self._gauge_cls or _PromGauge
Histogram = self._histogram_cls or _PromHistogram Histogram = self._histogram_cls or _PromHistogram
self.labels = labels or {} self.labels = labels or {}
self.startup_time_seconds = Gauge(
name="sglang:startup_time_seconds",
documentation="Engine startup duration by phase in seconds.",
labelnames=[*labels.keys(), "phase"],
multiprocess_mode="mostrecent",
)
self.startup_cuda_graph_time_seconds = Gauge(
name="sglang:startup_cuda_graph_time_seconds",
documentation="CUDA graph capture duration by phase in seconds.",
labelnames=[*labels.keys(), "phase"],
multiprocess_mode="mostrecent",
)
self.prompt_tokens_total = Counter( self.prompt_tokens_total = Counter(
name="sglang:prompt_tokens_total", name="sglang:prompt_tokens_total",
documentation="Number of prefill tokens processed.", documentation="Number of prefill tokens processed.",
@@ -1649,6 +1685,24 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
buckets=bucket_e2e_request_latency, buckets=bucket_e2e_request_latency,
) )
def emit_startup_time(self, startup_time: Mapping[str, Any]) -> None:
for phase in (
"load_weight",
"kv_cache_allocation",
"scheduler_e2e",
"tokenizer_e2e",
):
self.startup_time_seconds.labels(
**self.labels,
phase=phase,
).set(float(startup_time[phase]))
for phase, duration in startup_time["cuda_graph"].items():
self.startup_cuda_graph_time_seconds.labels(
**self.labels,
phase=phase,
).set(float(duration))
def observe_one_finished_request( def observe_one_finished_request(
self, self,
labels: Dict[str, str], labels: Dict[str, str],
@@ -284,6 +284,7 @@ class RayTokenizerMetricsCollector(TokenizerMetricsCollector):
"""``TokenizerMetricsCollector`` that emits via Ray's metric system.""" """``TokenizerMetricsCollector`` that emits via Ray's metric system."""
_counter_cls = RayCounterWrapper _counter_cls = RayCounterWrapper
_gauge_cls = RayGaugeWrapper
_histogram_cls = RayHistogramWrapper _histogram_cls = RayHistogramWrapper
@@ -0,0 +1,69 @@
from __future__ import annotations
from collections.abc import Iterable, Mapping
from sglang.srt.model_executor.graph_memory_usage import (
empty_graph_time_usage,
merge_graph_time_usage,
)
def build_scheduler_startup_time(
*,
target_load_weight: float,
draft_load_weight: float,
kv_cache_allocation: float,
scheduler_e2e: float,
target_cuda_graph: Mapping[str, float] | None,
draft_cuda_graph: Mapping[str, float] | None,
) -> dict:
return {
"load_weight": target_load_weight + draft_load_weight,
"kv_cache_allocation": kv_cache_allocation,
"scheduler_e2e": scheduler_e2e,
"cuda_graph": merge_graph_time_usage(
target_cuda_graph,
draft_cuda_graph,
),
}
def aggregate_scheduler_startup_times(
startup_times: Iterable[Mapping | None],
) -> dict:
"""Return critical-path startup durations across scheduler ranks."""
result = {
"load_weight": 0.0,
"kv_cache_allocation": 0.0,
"scheduler_e2e": 0.0,
"cuda_graph": empty_graph_time_usage(),
}
for startup_time in startup_times:
if not startup_time:
continue
result["load_weight"] = max(
result["load_weight"], float(startup_time.get("load_weight", 0.0))
)
result["kv_cache_allocation"] = max(
result["kv_cache_allocation"],
float(startup_time.get("kv_cache_allocation", 0.0)),
)
result["scheduler_e2e"] = max(
result["scheduler_e2e"],
float(startup_time.get("scheduler_e2e", 0.0)),
)
for phase, duration in startup_time.get("cuda_graph", {}).items():
result["cuda_graph"][phase] = max(
result["cuda_graph"].get(phase, 0.0), float(duration)
)
return result
def build_engine_startup_time(
scheduler_startup_times: Iterable[Mapping | None],
*,
tokenizer_e2e: float,
) -> dict:
result = aggregate_scheduler_startup_times(scheduler_startup_times)
result["tokenizer_e2e"] = tokenizer_e2e
return result
@@ -24,6 +24,7 @@ import zmq
from sglang.srt.entrypoints.engine import _calculate_rank_ranges from sglang.srt.entrypoints.engine import _calculate_rank_ranges
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
from sglang.srt.managers.data_parallel_controller import DataParallelController from sglang.srt.managers.data_parallel_controller import DataParallelController
from sglang.srt.observability.startup_time import aggregate_scheduler_startup_times
from sglang.srt.ray.engine import ( from sglang.srt.ray.engine import (
_compute_world_size, _compute_world_size,
_create_scheduler_actor, _create_scheduler_actor,
@@ -60,6 +61,7 @@ class RayDataParallelController(DataParallelController):
self.rank0_node_ip = rank0_node_ip self.rank0_node_ip = rank0_node_ip
self.scheduler_actors: List = [] self.scheduler_actors: List = []
self.event_loop_refs: List = [] self.event_loop_refs: List = []
self.startup_time = None
# super().__init__ will call our overridden launch methods via MRO. # super().__init__ will call our overridden launch methods via MRO.
# Pass run_scheduler_process_func=None since we don't spawn mp.Process. # Pass run_scheduler_process_func=None since we don't spawn mp.Process.
@@ -270,6 +272,10 @@ class RayDataParallelController(DataParallelController):
if scheduler_infos: if scheduler_infos:
self.max_total_num_tokens = scheduler_infos[0]["max_total_num_tokens"] self.max_total_num_tokens = scheduler_infos[0]["max_total_num_tokens"]
self.max_req_input_len = scheduler_infos[0]["max_req_input_len"] self.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
self.startup_time = aggregate_scheduler_startup_times(
[self.startup_time]
+ [info.get("startup_time") for info in scheduler_infos]
)
# Start event loops (non-blocking — runs until actor is killed) # Start event loops (non-blocking — runs until actor is killed)
self.event_loop_refs.extend( self.event_loop_refs.extend(
+1
View File
@@ -484,6 +484,7 @@ class RayEngine(Engine):
{ {
"max_total_num_tokens": controller.max_total_num_tokens, "max_total_num_tokens": controller.max_total_num_tokens,
"max_req_input_len": controller.max_req_input_len, "max_req_input_len": controller.max_req_input_len,
"startup_time": controller.startup_time,
} }
] ]
+5
View File
@@ -8350,6 +8350,11 @@ class ServerArgs:
if hasattr(self, "model_config"): if hasattr(self, "model_config"):
return self.model_config return self.model_config
self.model_config = ModelConfig.from_server_args(self) self.model_config = ModelConfig.from_server_args(self)
if self.model_config.is_hybrid_swa:
logger.info(
"Hybrid SWA model detected. architectures=%s",
self.model_config.hf_config.architectures,
)
return self.model_config return self.model_config
def _resolved(self): def _resolved(self):
@@ -5,6 +5,10 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.model_executor.graph_memory_usage import (
merge_graph_memory_usage,
merge_graph_time_usage,
)
from sglang.srt.runtime_context import get_exec, get_schedule from sglang.srt.runtime_context import get_exec, get_schedule
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -21,6 +25,10 @@ class EagleDraftWorkerBase(ABC):
_topk1_parents_prealloc: Optional[torch.Tensor] = None _topk1_parents_prealloc: Optional[torch.Tensor] = None
_topk1_score_indices_prealloc: Optional[torch.Tensor] = None _topk1_score_indices_prealloc: Optional[torch.Tensor] = None
def __init__(self) -> None:
self._specialized_graph_memory_usage: dict[str, float] = {}
self._specialized_graph_time_usage: dict[str, float] = {}
@abstractmethod @abstractmethod
def draft(): def draft():
pass pass
@@ -35,6 +43,24 @@ class EagleDraftWorkerBase(ABC):
per-step runner list.""" per-step runner list."""
return [self.draft_runner] return [self.draft_runner]
@property
def graph_memory_usage(self) -> dict[str, float]:
return merge_graph_memory_usage(
*(runner.graph_memory_usage for runner in self.draft_runners),
self._specialized_graph_memory_usage,
)
@property
def graph_time_usage(self) -> dict[str, float]:
return merge_graph_time_usage(
*(runner.graph_time_usage for runner in self.draft_runners),
self._specialized_graph_time_usage,
)
@property
def weight_load_time(self) -> float:
return sum(runner.weight_load_time for runner in self.draft_runners)
def alloc_memory_pool(self, **kwargs): def alloc_memory_pool(self, **kwargs):
pass pass
@@ -85,6 +111,10 @@ class EagleDraftWorkerBase(ABC):
class BaseSpecWorker(ABC): class BaseSpecWorker(ABC):
def __init__(self) -> None:
self._additional_graph_memory_usage: dict[str, float] = {}
self._additional_graph_time_usage: dict[str, float] = {}
@property @property
def target_worker(self) -> TpModelWorker: def target_worker(self) -> TpModelWorker:
return self._target_worker return self._target_worker
@@ -95,6 +125,34 @@ class BaseSpecWorker(ABC):
# ngram has no draft worker at all (returns None via its override). # ngram has no draft worker at all (returns None via its override).
return self._draft_worker return self._draft_worker
@property
def graph_memory_usage(self) -> dict[str, float]:
if self.draft_worker is None:
draft_memory_usage = None
else:
draft_memory_usage = self.draft_worker.graph_memory_usage
return merge_graph_memory_usage(
draft_memory_usage,
self._additional_graph_memory_usage,
)
@property
def graph_time_usage(self) -> dict[str, float]:
if self.draft_worker is None:
draft_time_usage = None
else:
draft_time_usage = self.draft_worker.graph_time_usage
return merge_graph_time_usage(
draft_time_usage,
self._additional_graph_time_usage,
)
@property
def weight_load_time(self) -> float:
if self.draft_worker is None:
return 0.0
return self.draft_worker.weight_load_time
@property @property
def war_fastpath_runner(self): def war_fastpath_runner(self):
# The runner that runs the step's LAST shared-buffer-reading phase -- # The runner that runs the step's LAST shared-buffer-reading phase --
@@ -171,6 +171,8 @@ class DFlashWorkerV2(BaseSpecWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps self.ps = ps
@@ -77,6 +77,8 @@ class DSparkWorkerV2(BaseSpecWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps self.ps = ps
@@ -132,6 +132,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
@@ -367,10 +369,20 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.target_worker.device self.target_worker.device
](self) ](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id) after_mem = get_available_gpu_memory(self.device, self.gpu_id)
capture_time = time.perf_counter() - tic
self._specialized_graph_memory_usage["draft_decode"] = (
self._specialized_graph_memory_usage.get("draft_decode", 0.0)
+ before_mem
- after_mem
)
self._specialized_graph_time_usage["draft_decode"] = (
self._specialized_graph_time_usage.get("draft_decode", 0.0)
+ capture_time
)
log_info_on_rank0( log_info_on_rank0(
logger, logger,
"Capture draft decode CUDA graph end. " "Capture draft decode CUDA graph end. "
f"elapsed={time.perf_counter() - tic:.2f} s, " f"elapsed={capture_time:.2f} s, "
f"mem usage={(before_mem - after_mem):.2f} GB, " f"mem usage={(before_mem - after_mem):.2f} GB, "
f"avail mem={after_mem:.2f} GB.", f"avail mem={after_mem:.2f} GB.",
) )
@@ -452,10 +464,20 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# draft_extend is the step's last shared-buffer-reading phase; its # draft_extend is the step's last shared-buffer-reading phase; its
# read-done event is what the scheduler's WAR barrier waits on. # read-done event is what the scheduler's WAR barrier waits on.
after_mem = get_available_gpu_memory(self.device, self.gpu_id) after_mem = get_available_gpu_memory(self.device, self.gpu_id)
capture_time = time.perf_counter() - tic
self._specialized_graph_memory_usage["draft_extend"] = (
self._specialized_graph_memory_usage.get("draft_extend", 0.0)
+ before_mem
- after_mem
)
self._specialized_graph_time_usage["draft_extend"] = (
self._specialized_graph_time_usage.get("draft_extend", 0.0)
+ capture_time
)
log_info_on_rank0( log_info_on_rank0(
logger, logger,
"Capture draft extend CUDA graph end. " "Capture draft extend CUDA graph end. "
f"elapsed={time.perf_counter() - tic:.2f} s, " f"elapsed={capture_time:.2f} s, "
f"mem usage={(before_mem - after_mem):.2f} GB, " f"mem usage={(before_mem - after_mem):.2f} GB, "
f"avail mem={after_mem:.2f} GB.", f"avail mem={after_mem:.2f} GB.",
) )
@@ -990,6 +1012,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
# Parse arguments # Parse arguments
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
@@ -1305,12 +1329,29 @@ class EAGLEWorkerV2(BaseSpecWorker):
TargetGraphRunnerCls = ( TargetGraphRunnerCls = (
NPUGraphRunner if _is_npu else DecodeCudaGraphRunner NPUGraphRunner if _is_npu else DecodeCudaGraphRunner
) )
target_graph_before_mem = get_available_gpu_memory(
self.device, self.gpu_id
)
target_graph_tic = time.perf_counter()
target_graph_runner = TargetGraphRunnerCls( target_graph_runner = TargetGraphRunnerCls(
target_model_runner, target_model_runner,
attn_backend=target_attn_backend, attn_backend=target_attn_backend,
speculative_num_steps=speculative_num_steps, speculative_num_steps=speculative_num_steps,
speculative_num_draft_tokens=speculative_num_draft_tokens, speculative_num_draft_tokens=speculative_num_draft_tokens,
) )
target_graph_after_mem = get_available_gpu_memory(
self.device, self.gpu_id
)
target_graph_time = time.perf_counter() - target_graph_tic
self._additional_graph_memory_usage["target_verify"] = (
self._additional_graph_memory_usage.get("target_verify", 0.0)
+ target_graph_before_mem
- target_graph_after_mem
)
self._additional_graph_time_usage["target_verify"] = (
self._additional_graph_time_usage.get("target_verify", 0.0)
+ target_graph_time
)
state = SpecRuntimeState( state = SpecRuntimeState(
speculative_num_steps=speculative_num_steps, speculative_num_steps=speculative_num_steps,
@@ -22,6 +22,7 @@ start of the next draft.
from __future__ import annotations from __future__ import annotations
import logging import logging
import time
from dataclasses import replace from dataclasses import replace
from typing import Optional from typing import Optional
@@ -43,7 +44,7 @@ from sglang.srt.model_executor.forward_batch_info import (
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import EagleDraftWorkerBase from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.eagle_utils import ( from sglang.srt.speculative.eagle_utils import (
build_tree_kernel_efficient, build_tree_kernel_efficient,
organize_draft_results, organize_draft_results,
@@ -70,7 +71,7 @@ from sglang.srt.speculative.spec_utils import (
select_top_k_tokens, select_top_k_tokens,
spec_stage_span, spec_stage_span,
) )
from sglang.srt.utils import empty_context from sglang.srt.utils import empty_context, get_available_gpu_memory
from sglang.srt.utils.async_probe import ( from sglang.srt.utils.async_probe import (
maybe_detect_inf, maybe_detect_inf,
maybe_detect_nan, maybe_detect_nan,
@@ -96,6 +97,8 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
EagleDraftWorkerBase.__init__(self)
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = server_args.speculative_num_steps
@@ -125,8 +128,8 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
with ( with (
empty_context() empty_context()
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context(): ), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
# NOTE: call TpModelWorker.__init__ explicitly -- EagleDraftWorkerBase is # Both base classes own initialization, so initialize TpModelWorker
# an ABC with no __init__, so cooperative super() would be ambiguous. # explicitly after EagleDraftWorkerBase above.
TpModelWorker.__init__( TpModelWorker.__init__(
self, self,
server_args=server_args, server_args=server_args,
@@ -362,7 +365,20 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
) )
logger.info("Capture Frozen-KV MTP draft cuda graph begin.") logger.info("Capture Frozen-KV MTP draft cuda graph begin.")
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
self.cuda_graph_runner = FrozenKVMTPCudaGraphRunner(self) self.cuda_graph_runner = FrozenKVMTPCudaGraphRunner(self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
self._specialized_graph_memory_usage["draft_decode"] = (
self._specialized_graph_memory_usage.get("draft_decode", 0.0)
+ before_mem
- after_mem
)
self._specialized_graph_time_usage["draft_decode"] = (
self._specialized_graph_time_usage.get("draft_decode", 0.0)
+ time.perf_counter()
- tic
)
logger.info("Capture Frozen-KV MTP draft cuda graph end.") logger.info("Capture Frozen-KV MTP draft cuda graph end.")
def _select_last_extend_hidden( def _select_last_extend_hidden(
@@ -660,6 +676,8 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
BaseSpecWorker.__init__(self)
# NOTE: intentionally does NOT call EAGLEWorkerV2.__init__ -- that builds # NOTE: intentionally does NOT call EAGLEWorkerV2.__init__ -- that builds
# an EagleDraftWorker (with its own draft KV pool). The frozen draft owns # an EagleDraftWorker (with its own draft KV pool). The frozen draft owns
# no KV, so we mirror the relevant setup and build a FrozenKVMTPDraftWorker. # no KV, so we mirror the relevant setup and build a FrozenKVMTPDraftWorker.
@@ -15,6 +15,7 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import time
from dataclasses import replace from dataclasses import replace
from typing import TYPE_CHECKING, List from typing import TYPE_CHECKING, List
@@ -76,7 +77,12 @@ from sglang.srt.speculative.spec_utils import (
sample_draft_proposal, sample_draft_proposal,
select_top_k_tokens, select_top_k_tokens,
) )
from sglang.srt.utils import is_cpu, is_npu, require_gathered_buffer from sglang.srt.utils import (
get_available_gpu_memory,
is_cpu,
is_npu,
require_gathered_buffer,
)
from sglang.srt.utils.async_probe import ( from sglang.srt.utils.async_probe import (
maybe_detect_inf, maybe_detect_inf,
maybe_detect_nan, maybe_detect_nan,
@@ -106,6 +112,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
@@ -372,6 +380,9 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
if envs.SGLANG_DISABLE_DRAFT_EXTEND_CUDA_GRAPH.get(): if envs.SGLANG_DISABLE_DRAFT_EXTEND_CUDA_GRAPH.get():
return return
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
if not _is_npu: if not _is_npu:
# The single-CG runner replays with no Python between steps, so the # The single-CG runner replays with no Python between steps, so the
# attn backend must fully rebuild its per-step metadata as captured # attn backend must fully rebuild its per-step metadata as captured
@@ -403,6 +414,17 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.cuda_graph_runner_for_draft_extend = ( self.cuda_graph_runner_for_draft_extend = (
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner(self) MultiLayerEagleMultiStepDraftExtendNpuGraphRunner(self)
) )
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
self._specialized_graph_memory_usage["draft_extend"] = (
self._specialized_graph_memory_usage.get("draft_extend", 0.0)
+ before_mem
- after_mem
)
self._specialized_graph_time_usage["draft_extend"] = (
self._specialized_graph_time_usage.get("draft_extend", 0.0)
+ time.perf_counter()
- tic
)
def draft(self, batch: ScheduleBatch): def draft(self, batch: ScheduleBatch):
draft_input: EagleDraftInput = batch.spec_info draft_input: EagleDraftInput = batch.spec_info
@@ -894,6 +916,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
# Parse arguments # Parse arguments
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
@@ -84,6 +84,8 @@ class NGRAMWorker(BaseSpecWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
super().__init__()
self.server_args = server_args self.server_args = server_args
self.enable_overlap = not server_args.disable_overlap_schedule self.enable_overlap = not server_args.disable_overlap_schedule
self._target_worker = target_worker self._target_worker = target_worker
@@ -11,6 +11,10 @@ from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import ( from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController, AdaptiveController,
) )
from sglang.srt.speculative.base_spec_worker import (
BaseSpecWorker,
EagleDraftWorkerBase,
)
from sglang.srt.speculative.eagle_utils import default_tree_mask_mode from sglang.srt.speculative.eagle_utils import default_tree_mask_mode
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2 from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -35,6 +39,8 @@ class StandaloneDraftWorker(EagleDraftWorker):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
EagleDraftWorkerBase.__init__(self)
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
@@ -109,15 +115,17 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.init_lm_head() self.init_lm_head()
def init_attention_backends(self): def init_attention_backends(self):
with self.draft_tp_context( with (
self.draft_runner.tp_group self.draft_tp_context(self.draft_runner.tp_group),
), speculative_moe_backend_context(): speculative_moe_backend_context(),
):
super().init_attention_backends() super().init_attention_backends()
def init_cuda_graphs(self): def init_cuda_graphs(self):
with self.draft_tp_context( with (
self.draft_runner.tp_group self.draft_tp_context(self.draft_runner.tp_group),
), speculative_moe_backend_context(): speculative_moe_backend_context(),
):
super().init_cuda_graphs() super().init_cuda_graphs()
def init_lm_head(self): def init_lm_head(self):
@@ -137,6 +145,8 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
BaseSpecWorker.__init__(self)
# Parse arguments # Parse arguments
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
+41
View File
@@ -41,6 +41,7 @@ class TestSRTEndpoint(CustomTestCase):
# Extra server-launch env; subclasses override to run the same suite # Extra server-launch env; subclasses override to run the same suite
# against a different server flavor (e.g. SGLANG_RUST_SERVER=1). # against a different server flavor (e.g. SGLANG_RUST_SERVER=1).
env = {} env = {}
expect_startup_observability = True
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -562,6 +563,45 @@ class TestSRTEndpoint(CustomTestCase):
version = response_json["version"] version = response_json["version"]
self.assertIsInstance(version, str) self.assertIsInstance(version, str)
if not self.expect_startup_observability:
return
startup_time = response_json["startup_time"]
for phase in (
"load_weight",
"kv_cache_allocation",
"scheduler_e2e",
"tokenizer_e2e",
):
self.assertIsInstance(startup_time[phase], float)
self.assertGreater(startup_time[phase], 0)
graph_phases = {
"prefill",
"decode",
"target_verify",
"draft_prefill",
"draft_decode",
"draft_extend",
}
self.assertTrue(graph_phases.issubset(startup_time["cuda_graph"]))
for phase in graph_phases:
self.assertIsInstance(startup_time["cuda_graph"][phase], float)
self.assertGreaterEqual(startup_time["cuda_graph"][phase], 0)
self.assertGreater(startup_time["cuda_graph"]["decode"], 0)
memory_usage = response_json["internal_states"][0]["memory_usage"]
self.assertIsInstance(memory_usage["weight"], float)
self.assertIsInstance(memory_usage["kvcache"], float)
self.assertEqual(memory_usage["token_capacity"], max_total_num_tokens)
self.assertIsNone(memory_usage["token_capacity_swa"])
self.assertIsInstance(memory_usage["startup_available"], float)
self.assertGreater(memory_usage["startup_available"], 0)
self.assertTrue(graph_phases.issubset(memory_usage["graph"]))
for phase in graph_phases:
self.assertIsInstance(memory_usage["graph"][phase], float)
self.assertGreaterEqual(memory_usage["graph"][phase], 0)
def test_logit_bias(self): def test_logit_bias(self):
"""Test that a very high logit bias forces sampling of a specific token.""" """Test that a very high logit bias forces sampling of a specific token."""
# Choose a token ID to bias (using 5 as an example) # Choose a token ID to bias (using 5 as an example)
@@ -864,6 +904,7 @@ class TestTokenizeDetokenize(CustomTestCase):
) )
class TestRustServerEndpoint(TestSRTEndpoint): class TestRustServerEndpoint(TestSRTEndpoint):
env = {"SGLANG_RUST_SERVER": "1"} env = {"SGLANG_RUST_SERVER": "1"}
expect_startup_observability = False
_RUST_TODO = "not implemented by the embedded Rust server yet" _RUST_TODO = "not implemented by the embedded Rust server yet"
+41 -10
View File
@@ -29,6 +29,14 @@ register_cuda_ci(est_time=74, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=32, suite="stage-b-test-1-gpu-small-amd") register_amd_ci(est_time=32, suite="stage-b-test-1-gpu-small-amd")
_MODEL_NAME = "Qwen/Qwen3-0.6B" _MODEL_NAME = "Qwen/Qwen3-0.6B"
_GRAPH_PHASES = {
"prefill",
"decode",
"target_verify",
"draft_prefill",
"draft_decode",
"draft_extend",
}
class TestEnableMetrics(CustomTestCase): class TestEnableMetrics(CustomTestCase):
@@ -66,14 +74,6 @@ class TestEnableMetrics(CustomTestCase):
"sglang:dp_cooperation_realtime_tokens_total", "sglang:dp_cooperation_realtime_tokens_total",
{"mode": "decode"}, {"mode": "decode"},
), ),
(
"sglang:dp_cooperation_forward_execution_seconds_total",
{"category": "extend"},
),
(
"sglang:dp_cooperation_forward_execution_seconds_total",
{"category": "decode"},
),
] ]
_check_metrics_positive(self, metrics, metrics_to_check) _check_metrics_positive(self, metrics, metrics_to_check)
@@ -139,8 +139,8 @@ class TestEnableMetrics(CustomTestCase):
for _ in response.iter_lines(decode_unicode=False): for _ in response.iter_lines(decode_unicode=False):
pass pass
for i in range(2): for _ in range(3):
# Send the request twice to trigger cached token metrics # The third request returns to the first rank under DP round-robin.
response = requests.post( response = requests.post(
f"{DEFAULT_URL_FOR_TEST}/generate", f"{DEFAULT_URL_FOR_TEST}/generate",
json={ json={
@@ -187,6 +187,12 @@ class TestEnableMetrics(CustomTestCase):
"sglang:num_unique_running_routing_keys", "sglang:num_unique_running_routing_keys",
"sglang:routing_key_running_req_count", "sglang:routing_key_running_req_count",
"sglang:routing_key_all_req_count", "sglang:routing_key_all_req_count",
"sglang:weight_memory_usage_gb",
"sglang:kv_cache_memory_usage_gb",
"sglang:graph_memory_usage_gb",
"sglang:startup_available_gpu_memory_gb",
"sglang:startup_time_seconds",
"sglang:startup_cuda_graph_time_seconds",
] ]
mfu_metrics = [ mfu_metrics = [
"sglang:estimated_flops_per_gpu_total", "sglang:estimated_flops_per_gpu_total",
@@ -224,9 +230,34 @@ class TestEnableMetrics(CustomTestCase):
("sglang:forward_execution_seconds_total", {"category": "extend"}), ("sglang:forward_execution_seconds_total", {"category": "extend"}),
("sglang:forward_execution_seconds_total", {"category": "decode"}), ("sglang:forward_execution_seconds_total", {"category": "decode"}),
("sglang:process_cpu_seconds_total", {"component": "tokenizer"}), ("sglang:process_cpu_seconds_total", {"component": "tokenizer"}),
("sglang:weight_memory_usage_gb", {"model_name": _MODEL_NAME}),
("sglang:kv_cache_memory_usage_gb", {"model_name": _MODEL_NAME}),
(
"sglang:startup_available_gpu_memory_gb",
{"model_name": _MODEL_NAME},
),
("sglang:startup_time_seconds", {"phase": "load_weight"}),
("sglang:startup_time_seconds", {"phase": "kv_cache_allocation"}),
("sglang:startup_time_seconds", {"phase": "scheduler_e2e"}),
("sglang:startup_time_seconds", {"phase": "tokenizer_e2e"}),
("sglang:startup_cuda_graph_time_seconds", {"phase": "decode"}),
] ]
_check_metrics_positive(self, metrics, metrics_to_check) _check_metrics_positive(self, metrics, metrics_to_check)
for metric_name in (
"sglang:graph_memory_usage_gb",
"sglang:startup_cuda_graph_time_seconds",
):
phases = {
sample.labels.get("phase")
for sample in metrics[metric_name]
if sample.labels.get("model_name") == _MODEL_NAME
}
self.assertTrue(
_GRAPH_PHASES.issubset(phases),
f"{metric_name}: missing graph phases {_GRAPH_PHASES - phases}",
)
if expect_mfu_metrics: if expect_mfu_metrics:
# Estimated perf metrics may have multiple series (e.g., by rank). Ensure # Estimated perf metrics may have multiple series (e.g., by rank). Ensure
# that at least one series for this model has a positive accumulated value. # that at least one series for this model has a positive accumulated value.
@@ -57,6 +57,7 @@ def _call_server_info_with(
tokenizer_manager.server_args = server_args tokenizer_manager.server_args = server_args
tokenizer_manager.model_path = server_args.model_path tokenizer_manager.model_path = server_args.model_path
tokenizer_manager.served_model_name = server_args.served_model_name tokenizer_manager.served_model_name = server_args.served_model_name
tokenizer_manager.startup_time = None
tokenizer_manager._config_updates = ( tokenizer_manager._config_updates = (
[("test", dict(config_updates))] if config_updates else [] [("test", dict(config_updates))] if config_updates else []
) )
@@ -77,12 +77,12 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
"check_cuda_graph_backend", "check_cuda_graph_backend",
return_value=False, return_value=False,
): ):
runner = capture_prefill_graph( capture = capture_prefill_graph(
model_runner=model_runner, model_runner=model_runner,
eager_runner=eager_runner, eager_runner=eager_runner,
) )
self.assertIs(runner, eager_runner) self.assertIs(capture.runner, eager_runner)
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self): def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
model_runner = SimpleNamespace( model_runner = SimpleNamespace(
@@ -225,11 +225,11 @@ class TestCollectorSubclassWiring(TestRayWrapperBase):
self.assertIs(cls._histogram_cls, self.rw.RayHistogramWrapper) self.assertIs(cls._histogram_cls, self.rw.RayHistogramWrapper)
self.assertIs(cls._summary_cls, self.rw.RaySummaryWrapper) self.assertIs(cls._summary_cls, self.rw.RaySummaryWrapper)
def test_tokenizer_overrides_counter_histogram_only(self): def test_tokenizer_overrides_counter_gauge_histogram(self):
cls = self.rw.RayTokenizerMetricsCollector cls = self.rw.RayTokenizerMetricsCollector
self.assertIs(cls._counter_cls, self.rw.RayCounterWrapper) self.assertIs(cls._counter_cls, self.rw.RayCounterWrapper)
self.assertIs(cls._gauge_cls, self.rw.RayGaugeWrapper)
self.assertIs(cls._histogram_cls, self.rw.RayHistogramWrapper) self.assertIs(cls._histogram_cls, self.rw.RayHistogramWrapper)
self.assertIsNone(cls._gauge_cls)
self.assertIsNone(cls._summary_cls) self.assertIsNone(cls._summary_cls)
def test_storage_overrides_counter_histogram_only(self): def test_storage_overrides_counter_histogram_only(self):