Add a test case for crash dump (#15905)

This commit is contained in:
Lianmin Zheng
2025-12-26 19:39:28 -08:00
committed by GitHub
parent 4edee6954a
commit a8380ded71
7 changed files with 307 additions and 209 deletions
@@ -345,11 +345,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
placeholder_tokens_val=None,
retraction_counts=recv_obj.retraction_counts,
token_steps=recv_obj.token_steps,
load=recv_obj.load,
queue_time=recv_obj.queue_time,
forward_entry_time=recv_obj.forward_entry_time,
prefill_launch_delay=recv_obj.prefill_launch_delay,
prefill_launch_latency=recv_obj.prefill_launch_latency,
load=recv_obj.load,
prefill_finished_ts=recv_obj.prefill_finished_ts,
)
@@ -386,12 +386,12 @@ def run_detokenizer_process(
try:
manager = detokenizer_manager_class(server_args, port_args)
if server_args.tokenizer_worker_num > 1:
manager.multi_http_worker_event_loop()
else:
if server_args.tokenizer_worker_num == 1:
manager.event_loop()
else:
manager.multi_http_worker_event_loop()
except Exception:
manager.maybe_clear_socket_mapping()
traceback = get_exception_traceback()
logger.error(f"DetokenizerManager hit an exception: {traceback}")
manager.maybe_clear_socket_mapping()
parent_process.send_signal(signal.SIGQUIT)
@@ -797,6 +797,22 @@ class SchedulerOutputProcessorMixin:
else: # embedding or reward model
self.stream_output_embedding(reqs)
if envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.get() > 0:
self._trigger_crash_for_tests(
envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.get()
)
def _trigger_crash_for_tests(self, crash_threshold: int):
# Crash trigger: crash after stream_output is called N times
# This is used for testing purposes.
if not hasattr(self, "_test_stream_output_count"):
self._test_stream_output_count = 0
self._test_stream_output_count += 1
if self._test_stream_output_count >= crash_threshold:
raise RuntimeError(
f"Test crash after stream_output called {self._test_stream_output_count} times"
)
def stream_output_generation(
self: Scheduler,
reqs: List[Req],
+146 -201
View File
@@ -20,6 +20,7 @@ import logging
import os
import pickle
import signal
import socket
import sys
import threading
import time
@@ -311,8 +312,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
def init_running_status(self):
# Request states
self._chosen_loop = None
self.rid_to_state: Dict[str, ReqState] = {}
self.event_loop = None
self.asyncio_tasks = set()
# Health check
@@ -324,6 +325,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.current_load = 0
self.current_load_lock = asyncio.Lock()
# Session
self.session_futures = {} # session_id -> asyncio event
def init_request_logging_and_dumping(self):
# Request logging
self.request_logger = RequestLogger(
@@ -338,6 +342,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.dump_request_list: List[Tuple] = []
self.crash_dump_request_list: deque[Tuple] = deque()
self.crash_dump_performed = False # Flag to ensure dump is only called once
self.straggler_request_list: List[Tuple] = []
# Initialize performance metrics loggers with proper skip names
_, obj_skip_names, out_skip_names = self.request_logger.metadata
@@ -351,9 +356,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if self.server_args.checkpoint_engine_wait_weights_before_ready:
self.initial_weights_loaded = False
# Session
self.session_futures = {} # session_id -> asyncio event
# Weight updates
# The event to notify the weight sync is finished.
self.model_update_lock = RWLock()
@@ -453,14 +455,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self,
obj: Union[GenerateReqInput, EmbeddingReqInput],
request: Optional[fastapi.Request] = None,
trace_parent: Optional[str] = None,
traceparent: Optional[str] = None,
):
created_time = obj.received_time if obj.received_time else time.time()
self.auto_create_handle_loop()
obj.normalize_batch_and_arguments()
# Normalize the request
obj.normalize_batch_and_arguments()
if self.enable_trace:
self._trace_request_start(obj, created_time, request, trace_parent)
self._trace_request_start(obj, created_time, request, traceparent)
if self.server_args.language_only:
self._handle_epd_disaggregation_encode_request(obj)
if self.server_args.tokenizer_worker_num > 1:
@@ -476,6 +479,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if self.server_args.enable_lora and obj.lora_path:
await self._resolve_lora_path(obj)
# Tokenize the request and send it to the scheduler
if obj.is_single:
tokenized_obj = await self._tokenize_one_request(obj)
state = self._send_one_request(obj, tokenized_obj, created_time)
@@ -1379,19 +1383,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
return background_tasks
def auto_create_handle_loop(self):
if self._chosen_loop is not None:
current_loop = get_or_create_event_loop()
assert (
current_loop == self._chosen_loop
), f"Please ensure only one event loop is ever used with SGLang. Previous loop: {self._chosen_loop}, current loop: {current_loop}"
if self.event_loop is not None:
return
# Create and start the handle_loop task
loop = get_or_create_event_loop()
self._chosen_loop = loop
self.asyncio_tasks.add(
loop.create_task(print_exception_wrapper(self.handle_loop))
)
self.event_loop = loop
# We cannot add signal handler when the tokenizer manager is not in
@@ -1413,131 +1412,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
loop.create_task(print_exception_wrapper(self.sigterm_watchdog))
)
def dump_requests_before_crash(self):
if self.crash_dump_performed:
logger.info(
"SIGTERM/SIGQUIT/Exception triggered, but crash dump already performed, skipping."
)
return
if not self.crash_dump_folder:
return
logger.error(f"Dumping requests before crash. {self.crash_dump_folder=}")
self.crash_dump_performed = True
# Check if NFS directory is available
# expected_nfs_dir = "/" + self.crash_dump_folder.lstrip("/").split("/")[0]
# use_nfs_dir = os.path.isdir(expected_nfs_dir) and os.access(
# expected_nfs_dir, os.W_OK
# )
use_nfs_dir = False
if not use_nfs_dir:
logger.error(
f"Expected NFS directory is not available or writable. Uploading to GCS."
)
data_to_dump = []
if self.crash_dump_request_list:
data_to_dump.extend(self.crash_dump_request_list)
# Add unfinished requests from rid_to_state
unfinished_requests = []
for rid, state in self.rid_to_state.items():
if not state.finished:
unfinished_requests.append(
(
state.obj,
state.out_list[-1] if state.out_list else {},
state.created_time,
time.time(),
)
)
if unfinished_requests:
data_to_dump.extend(unfinished_requests)
if not data_to_dump:
return
object_name = f'crash_dump_{datetime.now().strftime("%Y-%m-%d_%H-%M-%S")}.pkl'
filename = os.path.join(
self.crash_dump_folder,
os.getenv("HOSTNAME", None),
object_name,
)
os.makedirs(os.path.dirname(filename), exist_ok=True)
# Include server_args in the dump
data_to_dump_with_server_args = {
"server_args": self.server_args,
"requests": data_to_dump,
}
with open(filename, "wb") as f:
pickle.dump(data_to_dump_with_server_args, f)
logger.error(
f"Dumped {len(self.crash_dump_request_list)} finished and {len(unfinished_requests)} unfinished requests before crash to {filename}"
)
def _upload_file_to_gcs(bucket_name, source_file_path, object_name):
from google.cloud import storage
client = storage.Client()
bucket = client.bucket(bucket_name)
blob = bucket.blob(object_name)
blob.upload_from_filename(source_file_path, if_generation_match=0)
logger.error(
f"Successfully uploaded {source_file_path} to gs://{bucket_name}/{object_name}"
)
if not use_nfs_dir:
_upload_file_to_gcs(
"sglang_crash_dump",
filename,
os.getenv("HOSTNAME", None) + "/" + object_name,
)
async def sigterm_watchdog(self):
while not self.gracefully_exit:
await asyncio.sleep(5)
# Drain requests
while True:
remain_num_req = len(self.rid_to_state)
remaining_rids = list(self.rid_to_state.keys())
if self.server_status == ServerStatus.UnHealthy:
# if health check failed, we should exit immediately
logger.error(
"Signal SIGTERM received while health check failed. Force exiting."
)
self.dump_requests_before_crash()
self.force_exit_handler()
break
elif get_bool_env_var("SGL_FORCE_SHUTDOWN"):
# if force shutdown flag set, exit immediately
logger.error(
"Signal SIGTERM received while force shutdown flag set. Force exiting."
)
self.force_exit_handler()
break
logger.info(
f"Gracefully exiting... Remaining number of requests {remain_num_req}. Remaining requests {remaining_rids=}."
)
if remain_num_req > 0:
await asyncio.sleep(5)
else:
self.dump_requests_before_crash()
break
kill_process_tree(os.getpid(), include_parent=True)
sys.exit(0)
def force_exit_handler(self):
"""Put some custom force exit logic here."""
pass
async def handle_loop(self):
"""The event loop that handles requests"""
while True:
@@ -1547,28 +1421,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
self.last_receive_tstamp = time.time()
self.watchdog.feed()
def _add_metric_if_present(
self,
recv_obj: Any,
attr_name: str,
meta_info: Dict[str, Any],
index: int,
) -> None:
"""Add a metric to meta_info if it exists and is not None.
Args:
recv_obj: The received object that may contain the metric attribute
attr_name: The name of the attribute to check
meta_info: The dictionary to add the metric to
index: The index to access the metric value in the attribute list
"""
if (
hasattr(recv_obj, attr_name)
and getattr(recv_obj, attr_name)
and getattr(recv_obj, attr_name)[index] is not None
):
meta_info[attr_name] = getattr(recv_obj, attr_name)[index]
def _handle_batch_output(
self,
recv_obj: Union[
@@ -1676,12 +1528,12 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
state.finished = recv_obj.finished_reasons[i] is not None
if state.finished:
if self.server_args.speculative_algorithm:
self._calculate_spec_decoding_metrics(meta_info, recv_obj, i)
state.finished_time = time.time()
state.finished_time_perf = time.perf_counter()
meta_info["e2e_latency"] = state.finished_time - state.created_time
if self.server_args.speculative_algorithm:
self._calculate_spec_decoding_metrics(meta_info, recv_obj, i)
if self.enable_metrics:
self._calculate_timing_metrics(meta_info, state, recv_obj, i)
@@ -1708,10 +1560,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
# BatchTokenIDOutput.
if (
self.server_args.dp_size > 1
and (
isinstance(recv_obj, BatchStrOutput)
or isinstance(recv_obj, BatchTokenIDOutput)
)
and isinstance(recv_obj, (BatchStrOutput, BatchTokenIDOutput))
and recv_obj.load is not None
):
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load])
@@ -1875,21 +1724,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
i: int,
) -> None:
"""Calculate speculative decoding metrics, such as acceptance rate and acceptance length metrics."""
meta_info["spec_accept_rate"] = 0.0
meta_info["spec_accept_length"] = 0
meta_info["spec_verify_ct"] = recv_obj.spec_verify_ct[i]
# The draft tokens per speculative step (excluding the target-sampled token).
num_guess_tokens = self.server_args.speculative_num_draft_tokens - 1
if (
recv_obj.spec_verify_ct[i] > 0
and num_guess_tokens is not None
and not isinstance(recv_obj, BatchEmbeddingOutput)
hasattr(recv_obj, "spec_verify_ct")
and recv_obj.spec_verify_ct[i] > 0
and hasattr(recv_obj, "spec_accepted_tokens")
# Checks that `spec_accepted_tokens[i]` will exist.
and len(recv_obj.spec_accepted_tokens) > i
):
# The draft tokens per speculative step (excluding the target-sampled token).
num_guess_tokens = self.server_args.speculative_num_draft_tokens - 1
total_draft_tokens = recv_obj.spec_verify_ct[i] * num_guess_tokens
accepted_tokens = recv_obj.spec_accepted_tokens[i]
@@ -1950,6 +1792,28 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
completion_tokens = recv_obj.completion_tokens[i]
meta_info["decode_throughput"] = completion_tokens / decode_time
def _add_metric_if_present(
self,
recv_obj: Any,
attr_name: str,
meta_info: Dict[str, Any],
index: int,
) -> None:
"""Add a metric to meta_info if it exists and is not None.
Args:
recv_obj: The received object that may contain the metric attribute
attr_name: The name of the attribute to check
meta_info: The dictionary to add the metric to
index: The index to access the metric value in the attribute list
"""
if (
hasattr(recv_obj, attr_name)
and getattr(recv_obj, attr_name)
and getattr(recv_obj, attr_name)[index] is not None
):
meta_info[attr_name] = getattr(recv_obj, attr_name)[index]
def collect_metrics(self, state: ReqState, recv_obj: BatchStrOutput, i: int):
completion_tokens = (
recv_obj.completion_tokens[i]
@@ -2056,6 +1920,107 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
asyncio.create_task(asyncio.to_thread(background_task))
def dump_requests_before_crash(
self, hostname: str = os.getenv("HOSTNAME", socket.gethostname())
):
if not self.crash_dump_folder:
return
if self.crash_dump_performed:
logger.info(
"SIGTERM/SIGQUIT/Exception triggered, but crash dump already performed, skipping."
)
return
else:
self.crash_dump_performed = True
logger.error(f"Dumping requests before crash. {self.crash_dump_folder=}")
# Add finished requests from crash_dump_request_list
data_to_dump = []
if self.crash_dump_request_list:
data_to_dump.extend(self.crash_dump_request_list)
# Add unfinished requests from rid_to_state
unfinished_requests = []
for rid, state in self.rid_to_state.items():
if not state.finished:
unfinished_requests.append(
(
state.obj,
state.out_list[-1] if state.out_list else {},
state.created_time,
time.time(),
)
)
if unfinished_requests:
data_to_dump.extend(unfinished_requests)
if not data_to_dump:
return
# Create a file
filename = os.path.join(
self.crash_dump_folder,
hostname,
f'crash_dump_{datetime.now().strftime("%Y-%m-%d_%H-%M-%S")}.pkl',
)
os.makedirs(os.path.dirname(filename), exist_ok=True)
# Write the data to the file
data_to_dump_with_server_args = {
"server_args": self.server_args, # Include server_args in the dump
"requests": data_to_dump,
}
with open(filename, "wb") as f:
pickle.dump(data_to_dump_with_server_args, f)
logger.error(
f"Dumped {len(self.crash_dump_request_list)} finished and {len(unfinished_requests)} unfinished requests before crash to {filename}"
)
return filename
async def sigterm_watchdog(self):
while not self.gracefully_exit:
await asyncio.sleep(5)
# Drain requests
while True:
remain_num_req = len(self.rid_to_state)
remaining_rids = list(self.rid_to_state.keys())
if self.server_status == ServerStatus.UnHealthy:
# if health check failed, we should exit immediately
logger.error(
"Signal SIGTERM received while health check failed. Force exiting."
)
self.dump_requests_before_crash()
self.force_exit_handler()
break
elif get_bool_env_var("SGL_FORCE_SHUTDOWN"):
# if force shutdown flag set, exit immediately
logger.error(
"Signal SIGTERM received while force shutdown flag set. Force exiting."
)
self.force_exit_handler()
break
logger.info(
f"Gracefully exiting... Remaining number of requests {remain_num_req}. Remaining requests {remaining_rids=}."
)
if remain_num_req > 0:
await asyncio.sleep(5)
else:
self.dump_requests_before_crash()
break
kill_process_tree(os.getpid(), include_parent=True)
sys.exit(0)
def force_exit_handler(self):
"""Put some custom force exit logic here."""
pass
def _handle_abort_req(self, recv_obj: AbortReq):
if is_health_check_generate_req(recv_obj):
return
@@ -2107,26 +2072,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if len(self.model_update_tmp) == self.server_args.dp_size:
self.model_update_result.set_result(self.model_update_tmp)
def _extract_logprobs_for_tokens(
self, logprobs_data: List, label_token_ids: List[int]
) -> Dict[int, float]:
"""
Extract logprobs for specified token IDs from logprobs data.
Args:
logprobs_data: List of (logprob, token_id, text) tuples
label_token_ids: Token IDs to extract logprobs for
Returns:
Dictionary mapping token_id to logprob
"""
logprobs = {}
if logprobs_data:
for logprob, token_id, _ in logprobs_data:
if token_id in label_token_ids:
logprobs[token_id] = logprob
return logprobs
async def _resolve_lora_path(self, obj: Union[GenerateReqInput, EmbeddingReqInput]):
if isinstance(obj.lora_path, str):
unique_lora_paths = set([obj.lora_path])
@@ -2181,7 +2126,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
obj: Union[GenerateReqInput, EmbeddingReqInput],
created_time: Optional[float] = None,
request: Optional[fastapi.Request] = None,
trace_parent: Optional[str] = None,
traceparent: Optional[str] = None,
):
external_trace_header = None
if request:
@@ -2189,11 +2134,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
trace_set_remote_propagate_context(request.headers["trace_context"])
else:
external_trace_header = extract_trace_headers(request.headers)
elif trace_parent:
elif traceparent:
# When the request comes form the rust grpc server there isn't a
# real request object but we still need to propagate the traceparent from
# the traceparent that is explicitly passed in
external_trace_header = {"trace_parent": trace_parent}
external_trace_header = {"traceparent": traceparent}
if obj.is_single:
bootstrap_room = (
@@ -309,3 +309,23 @@ class TokenizerManagerMultiItemMixin:
]
return score_list
def _extract_logprobs_for_tokens(
self, logprobs_data: List, label_token_ids: List[int]
) -> Dict[int, float]:
"""
Extract logprobs for specified token IDs from logprobs data.
Args:
logprobs_data: List of (logprob, token_id, text) tuples
label_token_ids: Token IDs to extract logprobs for
Returns:
Dictionary mapping token_id to logprob
"""
logprobs = {}
if logprobs_data:
for logprob, token_id, _ in logprobs_data:
if token_id in label_token_ids:
logprobs[token_id] = logprob
return logprobs