[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:
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(
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.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.parser.template_detection import resolve_auto_parsers
from sglang.srt.parser.template_manager import TemplateManager
@@ -990,6 +991,35 @@ class Engine(EngineScoreMixin, EngineBase):
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
def _launch_subprocesses(
cls,
@@ -1010,6 +1040,8 @@ class Engine(EngineScoreMixin, EngineBase):
Returns:
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_logger(server_args)
_set_envs_and_config(server_args)
@@ -1144,6 +1176,8 @@ class Engine(EngineScoreMixin, EngineBase):
# Wait for the model to finish loading
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
tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[0][
"max_req_input_len"
@@ -1275,6 +1309,7 @@ class Engine(EngineScoreMixin, EngineBase):
dataclasses.asdict(self.tokenizer_manager.server_args)
),
**self._scheduler_init_result.scheduler_infos[0],
"startup_time": self.tokenizer_manager.startup_time,
"internal_states": internal_states,
"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.set_startup_time(scheduler_info["startup_time"])
set_global_state(
_GlobalState(
@@ -794,6 +795,7 @@ async def server_info():
dataclasses.asdict(server_args)
),
**_global_state.scheduler_info,
"startup_time": _global_state.tokenizer_manager.startup_time,
"internal_states": internal_states,
"version": __version__,
# 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.
app.is_single_tokenizer_mode = False
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:
@@ -294,7 +294,6 @@ class MlxModelRunnerStub(ModelRunner):
# No CUDA graphs, no attention backend
self.decode_cuda_graph_runner = None
self.graph_mem_usage = 0
self.attn_backend = None
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.observability.cpu_monitor import start_cpu_monitor_thread
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.runtime_context import get_exec, publish
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_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):
if req.routed_dp_rank is not None:
@@ -844,6 +848,7 @@ def run_data_parallel_controller_process(
"status": "ready",
"max_total_num_tokens": controller.max_total_num_tokens,
"max_req_input_len": controller.max_req_input_len,
"startup_time": controller.startup_time,
SCHEDULER_PIDS_ARG: scheduler_pids,
}
)
@@ -440,6 +440,7 @@ class MultiTokenizerRouter:
port_args: PortArgs,
):
self.server_args = server_args
self.startup_time: Optional[Dict[str, Any]] = None
context = zmq.asyncio.Context(3)
self.recv_from_detokenizer = get_zmq_socket(
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)
self.socket_mapping = SocketMapping()
def set_startup_time(self, startup_time: Dict[str, Any]) -> None:
self.startup_time = startup_time
def _run_loop(self):
self._loop.run_forever()
@@ -792,7 +796,9 @@ def read_from_shared_memory(name: str) -> Any:
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"""
# get main process ID
+98 -39
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 (
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 (
RECORD_STEP_TIME,
PrefillStats,
@@ -262,6 +266,7 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_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.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform
@@ -378,6 +383,11 @@ class Scheduler(
moe_dp_rank: 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
# init_soft_watchdog starts a daemon thread that reads these on its first tick.
self.forward_ct: int = 0
@@ -521,22 +531,8 @@ class Scheduler(
self.token_to_kv_pool_allocator = result.token_to_kv_pool_allocator
self.disable_radix_cache = result.disable_radix_cache
self.tree_cache = result.tree_cache
if _is_npu and is_deepseek_v4(
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)
self.emit_metrics_constants()
self.maybe_init_hccl_dp_prewarm()
if (c := self.tp_worker.model_runner.canary_manager) is not None:
c.attach_radix_cache(self.tree_cache)
@@ -645,6 +641,46 @@ class Scheduler(
self.init_batch_result_processor()
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):
if _is_npu:
@@ -941,14 +977,18 @@ class Scheduler(
self.maybe_init_draft_worker()
# Prepare KV cache pools for all workers
tic = time.perf_counter()
self.init_memory_pools()
self.kv_cache_allocation_time = time.perf_counter() - tic
self.init_all_attention_backends()
self.init_all_cuda_graphs()
model_runner = self.tp_worker.model_runner
if model_runner.token_to_kv_pool.post_capture_active:
tic = time.perf_counter()
model_runner.post_capture_resize_kv_pool()
self.kv_cache_allocation_time += time.perf_counter() - tic
if (
get_exec().moe.elastic_ep_backend is not None
@@ -1021,7 +1061,7 @@ class Scheduler(
set_random_seed(self.random_seed)
# 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
)
if self.ps.tp_rank == 0:
@@ -1031,23 +1071,36 @@ class Scheduler(
f"max_prefill_tokens={self.max_prefill_tokens}, "
f"max_running_requests={self.max_running_requests}, "
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:
self.metrics_collector.emit_constants(
max_total_num_tokens=self.max_total_num_tokens,
# TODO: max_running_requests_under_SLO has no setter — dead chain.
max_running_requests_under_SLO=getattr(
self, "max_running_requests_under_SLO", None
def emit_metrics_constants(self) -> None:
if not get_observability().enable_metrics:
return
self.metrics_collector.emit_constants(
max_total_num_tokens=self.max_total_num_tokens,
max_total_num_tokens_swa=self.swa_tokens_per_layer,
weight_memory_usage_gb=self.tp_worker.model_runner.weight_load_mem_usage,
kv_cache_memory_usage_gb=(
self.token_to_kv_pool_allocator.get_kvcache().mem_usage
),
graph_memory_usage_gb=combine_graph_memory_usage(
self.tp_worker.graph_memory_usage,
(
None
if self.draft_worker is None
else self.draft_worker.graph_memory_usage
),
engine_startup_time=0.0,
engine_load_weights_time=0.0,
page_size=self.page_size,
num_pages=self.max_total_num_tokens // self.page_size,
context_len=self.model_config.context_len,
startup_available_gpu_memory_gb=avail_mem,
)
),
# TODO: max_running_requests_under_SLO has no setter — dead chain.
max_running_requests_under_SLO=None,
page_size=self.page_size,
num_pages=self.max_total_num_tokens // self.page_size,
context_len=self.model_config.context_len,
startup_available_gpu_memory_gb=self.startup_available_gpu_memory_gb,
)
def init_hisparse_coordinator(self) -> None:
self.hisparse_coordinator: Optional[HiSparseCoordinator] = None
@@ -1561,6 +1614,7 @@ class Scheduler(
"status": "ready",
"max_total_num_tokens": self.max_total_num_tokens,
"max_req_input_len": self.max_req_input_len,
"startup_time": self.startup_time,
}
return result_dict
@@ -4113,14 +4167,19 @@ class Scheduler(
# readback reflects values changed via /set_internal_state, not startup.
ret = get_context().resolved_server_args_dict()
ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput
ret["memory_usage"] = {
"weight": round(self.tp_worker.model_runner.weight_load_mem_usage, 2),
"kvcache": round(
self.token_to_kv_pool_allocator.get_kvcache().mem_usage, 2
),
"token_capacity": int(self.max_total_num_tokens),
"graph": round(self.tp_worker.model_runner.graph_mem_usage, 2),
}
draft_graph_memory_usage = (
None if self.draft_worker is None else self.draft_worker.graph_memory_usage
)
ret["memory_usage"] = build_memory_usage(
weight_gb=self.tp_worker.model_runner.weight_load_mem_usage,
kv_cache_gb=self.token_to_kv_pool_allocator.get_kvcache().mem_usage,
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
if get_exec().moe.elastic_ep_backend is not None:
@@ -136,7 +136,7 @@ class SchedulerLoadInquirer:
kv_cache_gb=round(
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),
)
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
self.server_args = server_args
self.startup_time: Optional[Dict[str, Any]] = None
self._config_updates: List[Tuple[str, Dict[str, Any]]] = []
self.elastic_worker_count = server_args.dp_size
self.elastic_pending_ep_size = None
@@ -701,6 +702,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
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):
self._result_dispatcher = TypeBasedDispatcher(
[
+21
View File
@@ -46,6 +46,10 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
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.runtime_context import get_exec, get_model, get_schedule, get_spec
from sglang.srt.server_args import ServerArgs
@@ -97,6 +101,23 @@ class BaseTpWorker(ABC):
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):
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,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
allocation_label: Optional[str] = None,
):
self.size = size
self.page_size = page_size
@@ -1606,6 +1607,7 @@ class KVCache(abc.ABC):
self.layer_num = layer_num
self.start_layer = start_layer or 0
self.end_layer = end_layer or layer_num - 1
self.allocation_label = allocation_label
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
)
@@ -1626,19 +1628,27 @@ class KVCache(abc.ABC):
"""Common logging and mem_usage computation for KV cache allocation.
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()
if isinstance(kv_size_bytes, tuple):
k_size, v_size = kv_size_bytes
k_size_GB = k_size / GB
v_size_GB = v_size / GB
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
else:
kv_size_GB = kv_size_bytes / GB
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
@@ -1721,6 +1731,7 @@ class MHATokenToKVPool(KVCache):
kv_cache_layout: Optional[str] = None,
quant_method=None,
post_capture_active: bool = False,
allocation_label: Optional[str] = None,
):
self.k_buffer = None
self.v_buffer = None
@@ -1737,6 +1748,7 @@ class MHATokenToKVPool(KVCache):
enable_memory_saver,
start_layer,
end_layer,
allocation_label,
)
self.post_capture_active = post_capture_active
self._post_capture_owner = None
+12 -9
View File
@@ -57,19 +57,22 @@ class SWAKVPool(BaseSWAKVPool):
maybe_init_custom_mem_pool(device=self.device)
)
self.swa_kv_pool = token_to_kv_pool_class(
size=size_swa,
dtype=dtype,
layer_num=self.swa_layer_nums,
**kwargs,
)
kwargs.pop("swa_head_num", None)
kwargs.pop("swa_head_dim", None)
kwargs.pop("swa_v_head_dim", None)
full_pool_kwargs = kwargs.copy()
full_pool_kwargs.pop("swa_head_num", None)
full_pool_kwargs.pop("swa_head_dim", None)
full_pool_kwargs.pop("swa_v_head_dim", None)
self.full_kv_pool = token_to_kv_pool_class(
size=size,
dtype=dtype,
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,
)
# {layer_id: (index, is_swa_layer)}
@@ -1265,11 +1265,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
dp_padding_mode = DpPaddingMode.SUM_LEN
# 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
# count and can select a different capture bucket, so the in-graph DP
# collectives (all_gather / reduce_scatter) mismatch across ranks and
# corrupt the output. Force MAX_LEN so every rank pads to the global
# max and picks the same bucket (mirrors the decode cuda graph
# contract, which always runs MAX_LEN).
# count and can select a different capture bucket. This mismatches the
# rank-coupled communication geometry: DP gather/combine uses
# all_gather_into_tensor / reduce_scatter_tensor, while MoE backends may
# use A2A dispatch/combine. Force MAX_LEN so every rank pads to the global
# max and picks the same bucket.
#
# Only force MAX_LEN when the batch fits a captured breakable prefill
# 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,
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.attention_backend_setup import (
build_attention_backends,
@@ -310,6 +314,8 @@ class ModelRunner:
self.draft_model_idx = draft_model_idx
self.enable_hisparse = server_args.enable_hisparse
self.init_startup_observability()
self.init_remote_instance_weight_transporter()
self.init_msprobe()
@@ -411,6 +417,11 @@ class ModelRunner:
self.init_weight_updater()
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:
if not (
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
)
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.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):
if self.is_draft_worker:
@@ -1068,13 +1080,14 @@ class ModelRunner:
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_time = time.perf_counter() - tic_total
# Get quantization config from ModelConfig
# This handles both config.json (standard) and hf_quant_config.json (ModelOpt)
quant_str = self.model_config.get_quantization_config_log_str()
logger.info(
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"{quant_str + ', ' if quant_str else ''}"
f"avail mem={after_avail_memory:.2f} GB, "
@@ -1212,18 +1225,37 @@ class ModelRunner:
def init_decode_cuda_graph(self):
self.decode_cuda_graph_runner = None
self.graph_mem_usage = 0
capture = capture_decode_graph(model_runner=self)
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):
self.prefill_cuda_graph_runner = None
self.prefill_cuda_graph_runner = capture_prefill_graph(
capture = capture_prefill_graph(
model_runner=self,
eager_runner=self.eager_runner,
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):
self.local_omp_cpuid = numa_utils.init_threads_binding(
@@ -23,6 +23,10 @@ from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
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.hook_manager import register_forward_hooks
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]
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):
eager_runner: EagerRunner
prefill_runner: Optional[BaseRunner]
decode: DecodeGraphCapture
prefill: GraphCapture
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(
@@ -100,11 +128,17 @@ def capture_cuda_graphs(
# cuda-graph capture: prefill before decode, so both coalesce onto the
# eager buffer allocated above. (capture_prefill_graph routes prefill
# 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
)
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 model_runner.device in ("cuda", "musa", "cpu", "npu", "xpu"):
decode = capture_decode_graph(model_runner=model_runner)
@@ -113,7 +147,12 @@ def capture_cuda_graphs(
):
decode = capture_decode_graph(model_runner=model_runner)
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
# 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:
model_runner.canary_manager.mark_init_finished()
return CudaGraphsCapture(
eager_runner=eager_runner, prefill_runner=prefill_runner, decode=decode
)
return CudaGraphsCapture(eager_runner=eager_runner, prefill=prefill, decode=decode)
def capture_prefill_graph(
@@ -144,8 +181,22 @@ def capture_prefill_graph(
model_runner: ModelRunner,
eager_runner: EagerRunner,
force_for_draft_worker: bool = False,
) -> Optional[BaseRunner]:
"""Initialize prefill CUDA graph runner."""
) -> GraphCapture:
"""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):
logger.info(
@@ -157,14 +208,14 @@ def capture_prefill_graph(
# EagerRunner (its can_run_graph returns False, so _forward_raw's
# extend branch falls through to the eager path).
if not model_runner.is_draft_worker:
return eager_runner
return None
return result(eager_runner)
return result(None)
# Draft models skip here during __init__; the eagle worker calls
# this method explicitly (force_for_draft_worker=True) after
# init_lm_head so graphs capture the final embedding weights.
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
# 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 "
"to avoid FP4/MoE decode-replay corruption (#28386)."
)
return eager_runner
return result(eager_runner)
if (
model_runner.server_args.enable_lora
@@ -194,7 +245,7 @@ def capture_prefill_graph(
"configuration does not support it (unsupported LoRA backend, "
"MoE LoRA, or DP attention)."
)
return eager_runner
return result(eager_runner)
# Resolve the decoder once. Some VLM wrappers (for example Kimi-VL)
# expose it as ``language_model`` rather than ``model``.
@@ -204,12 +255,12 @@ def capture_prefill_graph(
logger.warning(
"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
if not model_runner.server_args.cuda_graph_config.prefill.bs:
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_backend = prefill_config.backend
@@ -272,7 +323,7 @@ def capture_prefill_graph(
logger.warning(
"Disable prefill CUDA graph because the model does not have a 'layers' attribute"
)
return None
return result(None)
(
model_runner.attention_layers,
@@ -288,7 +339,7 @@ def capture_prefill_graph(
logger,
"Disable prefill CUDA graph because some layers do not apply Standard GQA",
)
return None
return result(None)
tic = time.perf_counter()
before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
@@ -304,7 +355,7 @@ def capture_prefill_graph(
before_mem,
_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"
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)
mem_usage = before_mem - after_mem
capture_time = time.perf_counter() - tic
logger.info(
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."
)
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."""
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:
# 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)
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(
f"Capture {capture_name} {graph_backend[model_runner.device]} end. "
f"elapsed={time.perf_counter() - tic:.2f} s, "
f"mem usage={graph_mem_usage:.2f} GB, avail mem={after_mem:.2f} GB."
f"elapsed={capture_time:.2f} s, "
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
from collections import Counter
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.environ import envs
@@ -995,24 +995,36 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labelnames=labels.keys(),
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(
name="sglang:max_running_requests_under_SLO",
documentation="The maximum number of running requests under SLO.",
labelnames=labels.keys(),
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(
name="sglang:page_size",
documentation="KV cache page size in tokens.",
@@ -1401,21 +1413,30 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
def emit_constants(
self,
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],
engine_startup_time: float,
engine_load_weights_time: float,
page_size: int,
num_pages: int,
context_len: int,
startup_available_gpu_memory_gb: float,
) -> None:
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:
self._log_gauge(
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.num_pages, num_pages)
self._log_gauge(self.context_len, context_len)
@@ -1435,13 +1456,28 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
) -> None:
# 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 Gauge as _PromGauge
from prometheus_client import Histogram as _PromHistogram
Counter = self._counter_cls or _PromCounter
Gauge = self._gauge_cls or _PromGauge
Histogram = self._histogram_cls or _PromHistogram
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(
name="sglang:prompt_tokens_total",
documentation="Number of prefill tokens processed.",
@@ -1649,6 +1685,24 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
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(
self,
labels: Dict[str, str],
@@ -284,6 +284,7 @@ class RayTokenizerMetricsCollector(TokenizerMetricsCollector):
"""``TokenizerMetricsCollector`` that emits via Ray's metric system."""
_counter_cls = RayCounterWrapper
_gauge_cls = RayGaugeWrapper
_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.layers.dp_attention import compute_dp_attention_world_info
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 (
_compute_world_size,
_create_scheduler_actor,
@@ -60,6 +61,7 @@ class RayDataParallelController(DataParallelController):
self.rank0_node_ip = rank0_node_ip
self.scheduler_actors: List = []
self.event_loop_refs: List = []
self.startup_time = None
# super().__init__ will call our overridden launch methods via MRO.
# Pass run_scheduler_process_func=None since we don't spawn mp.Process.
@@ -270,6 +272,10 @@ class RayDataParallelController(DataParallelController):
if scheduler_infos:
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.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)
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_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"):
return self.model_config
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
def _resolved(self):
@@ -5,6 +5,10 @@ from typing import TYPE_CHECKING, Optional
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
if TYPE_CHECKING:
@@ -21,6 +25,10 @@ class EagleDraftWorkerBase(ABC):
_topk1_parents_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
def draft():
pass
@@ -35,6 +43,24 @@ class EagleDraftWorkerBase(ABC):
per-step runner list."""
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):
pass
@@ -85,6 +111,10 @@ class EagleDraftWorkerBase(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
def target_worker(self) -> TpModelWorker:
return self._target_worker
@@ -95,6 +125,34 @@ class BaseSpecWorker(ABC):
# ngram has no draft worker at all (returns None via its override).
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
def war_fastpath_runner(self):
# The runner that runs the step's LAST shared-buffer-reading phase --
@@ -171,6 +171,8 @@ class DFlashWorkerV2(BaseSpecWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
@@ -77,6 +77,8 @@ class DSparkWorkerV2(BaseSpecWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
@@ -132,6 +132,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
@@ -367,10 +369,20 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.target_worker.device
](self)
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(
logger,
"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"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
# read-done event is what the scheduler's WAR barrier waits on.
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(
logger,
"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"avail mem={after_mem:.2f} GB.",
)
@@ -990,6 +1012,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
@@ -1305,12 +1329,29 @@ class EAGLEWorkerV2(BaseSpecWorker):
TargetGraphRunnerCls = (
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_model_runner,
attn_backend=target_attn_backend,
speculative_num_steps=speculative_num_steps,
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(
speculative_num_steps=speculative_num_steps,
@@ -22,6 +22,7 @@ start of the next draft.
from __future__ import annotations
import logging
import time
from dataclasses import replace
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.pool_configurator import MemoryPoolConfig
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 (
build_tree_kernel_efficient,
organize_draft_results,
@@ -70,7 +71,7 @@ from sglang.srt.speculative.spec_utils import (
select_top_k_tokens,
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 (
maybe_detect_inf,
maybe_detect_nan,
@@ -96,6 +97,8 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
EagleDraftWorkerBase.__init__(self)
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
@@ -125,8 +128,8 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
with (
empty_context()
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
# NOTE: call TpModelWorker.__init__ explicitly -- EagleDraftWorkerBase is
# an ABC with no __init__, so cooperative super() would be ambiguous.
# Both base classes own initialization, so initialize TpModelWorker
# explicitly after EagleDraftWorkerBase above.
TpModelWorker.__init__(
self,
server_args=server_args,
@@ -362,7 +365,20 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
)
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)
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.")
def _select_last_extend_hidden(
@@ -660,6 +676,8 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
nccl_port: int,
target_worker: TpModelWorker,
):
BaseSpecWorker.__init__(self)
# NOTE: intentionally does NOT call EAGLEWorkerV2.__init__ -- that builds
# an EagleDraftWorker (with its own draft KV pool). The frozen draft owns
# no KV, so we mirror the relevant setup and build a FrozenKVMTPDraftWorker.
@@ -15,6 +15,7 @@
from __future__ import annotations
import logging
import time
from dataclasses import replace
from typing import TYPE_CHECKING, List
@@ -76,7 +77,12 @@ from sglang.srt.speculative.spec_utils import (
sample_draft_proposal,
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 (
maybe_detect_inf,
maybe_detect_nan,
@@ -106,6 +112,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
@@ -372,6 +380,9 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
if envs.SGLANG_DISABLE_DRAFT_EXTEND_CUDA_GRAPH.get():
return
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
if not _is_npu:
# The single-CG runner replays with no Python between steps, so the
# 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 = (
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):
draft_input: EagleDraftInput = batch.spec_info
@@ -894,6 +916,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
@@ -84,6 +84,8 @@ class NGRAMWorker(BaseSpecWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
self.server_args = server_args
self.enable_overlap = not server_args.disable_overlap_schedule
self._target_worker = target_worker
@@ -11,6 +11,10 @@ from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import (
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_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -35,6 +39,8 @@ class StandaloneDraftWorker(EagleDraftWorker):
nccl_port: int,
target_worker: TpModelWorker,
):
EagleDraftWorkerBase.__init__(self)
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
@@ -109,15 +115,17 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.init_lm_head()
def init_attention_backends(self):
with self.draft_tp_context(
self.draft_runner.tp_group
), speculative_moe_backend_context():
with (
self.draft_tp_context(self.draft_runner.tp_group),
speculative_moe_backend_context(),
):
super().init_attention_backends()
def init_cuda_graphs(self):
with self.draft_tp_context(
self.draft_runner.tp_group
), speculative_moe_backend_context():
with (
self.draft_tp_context(self.draft_runner.tp_group),
speculative_moe_backend_context(),
):
super().init_cuda_graphs()
def init_lm_head(self):
@@ -137,6 +145,8 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
nccl_port: int,
target_worker: TpModelWorker,
):
BaseSpecWorker.__init__(self)
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk