Add a test case for crash dump (#15905)
This commit is contained in:
@@ -151,6 +151,7 @@ class Envs:
|
|||||||
SGLANG_TEST_STUCK_DETOKENIZER = EnvFloat(0)
|
SGLANG_TEST_STUCK_DETOKENIZER = EnvFloat(0)
|
||||||
SGLANG_TEST_STUCK_DP_CONTROLLER = EnvFloat(0)
|
SGLANG_TEST_STUCK_DP_CONTROLLER = EnvFloat(0)
|
||||||
SGLANG_TEST_STUCK_TOKENIZER = EnvFloat(0)
|
SGLANG_TEST_STUCK_TOKENIZER = EnvFloat(0)
|
||||||
|
SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS = EnvInt(0)
|
||||||
IS_BLACKWELL = EnvBool(False)
|
IS_BLACKWELL = EnvBool(False)
|
||||||
IS_H200 = EnvBool(False)
|
IS_H200 = EnvBool(False)
|
||||||
SGLANG_SET_CPU_AFFINITY = EnvBool(False)
|
SGLANG_SET_CPU_AFFINITY = EnvBool(False)
|
||||||
|
|||||||
@@ -345,11 +345,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
placeholder_tokens_val=None,
|
placeholder_tokens_val=None,
|
||||||
retraction_counts=recv_obj.retraction_counts,
|
retraction_counts=recv_obj.retraction_counts,
|
||||||
token_steps=recv_obj.token_steps,
|
token_steps=recv_obj.token_steps,
|
||||||
|
load=recv_obj.load,
|
||||||
queue_time=recv_obj.queue_time,
|
queue_time=recv_obj.queue_time,
|
||||||
forward_entry_time=recv_obj.forward_entry_time,
|
forward_entry_time=recv_obj.forward_entry_time,
|
||||||
prefill_launch_delay=recv_obj.prefill_launch_delay,
|
prefill_launch_delay=recv_obj.prefill_launch_delay,
|
||||||
prefill_launch_latency=recv_obj.prefill_launch_latency,
|
prefill_launch_latency=recv_obj.prefill_launch_latency,
|
||||||
load=recv_obj.load,
|
|
||||||
prefill_finished_ts=recv_obj.prefill_finished_ts,
|
prefill_finished_ts=recv_obj.prefill_finished_ts,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -386,12 +386,12 @@ def run_detokenizer_process(
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
manager = detokenizer_manager_class(server_args, port_args)
|
manager = detokenizer_manager_class(server_args, port_args)
|
||||||
if server_args.tokenizer_worker_num > 1:
|
if server_args.tokenizer_worker_num == 1:
|
||||||
manager.multi_http_worker_event_loop()
|
|
||||||
else:
|
|
||||||
manager.event_loop()
|
manager.event_loop()
|
||||||
|
else:
|
||||||
|
manager.multi_http_worker_event_loop()
|
||||||
except Exception:
|
except Exception:
|
||||||
manager.maybe_clear_socket_mapping()
|
|
||||||
traceback = get_exception_traceback()
|
traceback = get_exception_traceback()
|
||||||
logger.error(f"DetokenizerManager hit an exception: {traceback}")
|
logger.error(f"DetokenizerManager hit an exception: {traceback}")
|
||||||
|
manager.maybe_clear_socket_mapping()
|
||||||
parent_process.send_signal(signal.SIGQUIT)
|
parent_process.send_signal(signal.SIGQUIT)
|
||||||
|
|||||||
@@ -797,6 +797,22 @@ class SchedulerOutputProcessorMixin:
|
|||||||
else: # embedding or reward model
|
else: # embedding or reward model
|
||||||
self.stream_output_embedding(reqs)
|
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(
|
def stream_output_generation(
|
||||||
self: Scheduler,
|
self: Scheduler,
|
||||||
reqs: List[Req],
|
reqs: List[Req],
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import pickle
|
import pickle
|
||||||
import signal
|
import signal
|
||||||
|
import socket
|
||||||
import sys
|
import sys
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
@@ -311,8 +312,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
|
|
||||||
def init_running_status(self):
|
def init_running_status(self):
|
||||||
# Request states
|
# Request states
|
||||||
self._chosen_loop = None
|
|
||||||
self.rid_to_state: Dict[str, ReqState] = {}
|
self.rid_to_state: Dict[str, ReqState] = {}
|
||||||
|
self.event_loop = None
|
||||||
self.asyncio_tasks = set()
|
self.asyncio_tasks = set()
|
||||||
|
|
||||||
# Health check
|
# Health check
|
||||||
@@ -324,6 +325,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
self.current_load = 0
|
self.current_load = 0
|
||||||
self.current_load_lock = asyncio.Lock()
|
self.current_load_lock = asyncio.Lock()
|
||||||
|
|
||||||
|
# Session
|
||||||
|
self.session_futures = {} # session_id -> asyncio event
|
||||||
|
|
||||||
def init_request_logging_and_dumping(self):
|
def init_request_logging_and_dumping(self):
|
||||||
# Request logging
|
# Request logging
|
||||||
self.request_logger = RequestLogger(
|
self.request_logger = RequestLogger(
|
||||||
@@ -338,6 +342,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
self.dump_request_list: List[Tuple] = []
|
self.dump_request_list: List[Tuple] = []
|
||||||
self.crash_dump_request_list: deque[Tuple] = deque()
|
self.crash_dump_request_list: deque[Tuple] = deque()
|
||||||
self.crash_dump_performed = False # Flag to ensure dump is only called once
|
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
|
# Initialize performance metrics loggers with proper skip names
|
||||||
_, obj_skip_names, out_skip_names = self.request_logger.metadata
|
_, 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:
|
if self.server_args.checkpoint_engine_wait_weights_before_ready:
|
||||||
self.initial_weights_loaded = False
|
self.initial_weights_loaded = False
|
||||||
|
|
||||||
# Session
|
|
||||||
self.session_futures = {} # session_id -> asyncio event
|
|
||||||
|
|
||||||
# Weight updates
|
# Weight updates
|
||||||
# The event to notify the weight sync is finished.
|
# The event to notify the weight sync is finished.
|
||||||
self.model_update_lock = RWLock()
|
self.model_update_lock = RWLock()
|
||||||
@@ -453,14 +455,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
self,
|
self,
|
||||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||||
request: Optional[fastapi.Request] = None,
|
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()
|
created_time = obj.received_time if obj.received_time else time.time()
|
||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
obj.normalize_batch_and_arguments()
|
|
||||||
|
|
||||||
|
# Normalize the request
|
||||||
|
obj.normalize_batch_and_arguments()
|
||||||
if self.enable_trace:
|
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:
|
if self.server_args.language_only:
|
||||||
self._handle_epd_disaggregation_encode_request(obj)
|
self._handle_epd_disaggregation_encode_request(obj)
|
||||||
if self.server_args.tokenizer_worker_num > 1:
|
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:
|
if self.server_args.enable_lora and obj.lora_path:
|
||||||
await self._resolve_lora_path(obj)
|
await self._resolve_lora_path(obj)
|
||||||
|
|
||||||
|
# Tokenize the request and send it to the scheduler
|
||||||
if obj.is_single:
|
if obj.is_single:
|
||||||
tokenized_obj = await self._tokenize_one_request(obj)
|
tokenized_obj = await self._tokenize_one_request(obj)
|
||||||
state = self._send_one_request(obj, tokenized_obj, created_time)
|
state = self._send_one_request(obj, tokenized_obj, created_time)
|
||||||
@@ -1379,19 +1383,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
return background_tasks
|
return background_tasks
|
||||||
|
|
||||||
def auto_create_handle_loop(self):
|
def auto_create_handle_loop(self):
|
||||||
if self._chosen_loop is not None:
|
if self.event_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}"
|
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Create and start the handle_loop task
|
||||||
loop = get_or_create_event_loop()
|
loop = get_or_create_event_loop()
|
||||||
self._chosen_loop = loop
|
|
||||||
self.asyncio_tasks.add(
|
self.asyncio_tasks.add(
|
||||||
loop.create_task(print_exception_wrapper(self.handle_loop))
|
loop.create_task(print_exception_wrapper(self.handle_loop))
|
||||||
)
|
)
|
||||||
|
|
||||||
self.event_loop = loop
|
self.event_loop = loop
|
||||||
|
|
||||||
# We cannot add signal handler when the tokenizer manager is not in
|
# 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))
|
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):
|
async def handle_loop(self):
|
||||||
"""The event loop that handles requests"""
|
"""The event loop that handles requests"""
|
||||||
while True:
|
while True:
|
||||||
@@ -1547,28 +1421,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
self.last_receive_tstamp = time.time()
|
self.last_receive_tstamp = time.time()
|
||||||
self.watchdog.feed()
|
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(
|
def _handle_batch_output(
|
||||||
self,
|
self,
|
||||||
recv_obj: Union[
|
recv_obj: Union[
|
||||||
@@ -1676,12 +1528,12 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
|
|
||||||
state.finished = recv_obj.finished_reasons[i] is not None
|
state.finished = recv_obj.finished_reasons[i] is not None
|
||||||
if state.finished:
|
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 = time.time()
|
||||||
state.finished_time_perf = time.perf_counter()
|
state.finished_time_perf = time.perf_counter()
|
||||||
meta_info["e2e_latency"] = state.finished_time - state.created_time
|
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:
|
if self.enable_metrics:
|
||||||
self._calculate_timing_metrics(meta_info, state, recv_obj, i)
|
self._calculate_timing_metrics(meta_info, state, recv_obj, i)
|
||||||
|
|
||||||
@@ -1708,10 +1560,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
# BatchTokenIDOutput.
|
# BatchTokenIDOutput.
|
||||||
if (
|
if (
|
||||||
self.server_args.dp_size > 1
|
self.server_args.dp_size > 1
|
||||||
and (
|
and isinstance(recv_obj, (BatchStrOutput, BatchTokenIDOutput))
|
||||||
isinstance(recv_obj, BatchStrOutput)
|
|
||||||
or isinstance(recv_obj, BatchTokenIDOutput)
|
|
||||||
)
|
|
||||||
and recv_obj.load is not None
|
and recv_obj.load is not None
|
||||||
):
|
):
|
||||||
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load])
|
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load])
|
||||||
@@ -1875,21 +1724,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
i: int,
|
i: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Calculate speculative decoding metrics, such as acceptance rate and acceptance length metrics."""
|
"""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 (
|
if (
|
||||||
recv_obj.spec_verify_ct[i] > 0
|
hasattr(recv_obj, "spec_verify_ct")
|
||||||
and num_guess_tokens is not None
|
and recv_obj.spec_verify_ct[i] > 0
|
||||||
and not isinstance(recv_obj, BatchEmbeddingOutput)
|
|
||||||
and hasattr(recv_obj, "spec_accepted_tokens")
|
and hasattr(recv_obj, "spec_accepted_tokens")
|
||||||
# Checks that `spec_accepted_tokens[i]` will exist.
|
|
||||||
and len(recv_obj.spec_accepted_tokens) > i
|
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
|
total_draft_tokens = recv_obj.spec_verify_ct[i] * num_guess_tokens
|
||||||
accepted_tokens = recv_obj.spec_accepted_tokens[i]
|
accepted_tokens = recv_obj.spec_accepted_tokens[i]
|
||||||
|
|
||||||
@@ -1950,6 +1792,28 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
completion_tokens = recv_obj.completion_tokens[i]
|
completion_tokens = recv_obj.completion_tokens[i]
|
||||||
meta_info["decode_throughput"] = completion_tokens / decode_time
|
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):
|
def collect_metrics(self, state: ReqState, recv_obj: BatchStrOutput, i: int):
|
||||||
completion_tokens = (
|
completion_tokens = (
|
||||||
recv_obj.completion_tokens[i]
|
recv_obj.completion_tokens[i]
|
||||||
@@ -2056,6 +1920,107 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
|
|
||||||
asyncio.create_task(asyncio.to_thread(background_task))
|
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):
|
def _handle_abort_req(self, recv_obj: AbortReq):
|
||||||
if is_health_check_generate_req(recv_obj):
|
if is_health_check_generate_req(recv_obj):
|
||||||
return
|
return
|
||||||
@@ -2107,26 +2072,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
if len(self.model_update_tmp) == self.server_args.dp_size:
|
if len(self.model_update_tmp) == self.server_args.dp_size:
|
||||||
self.model_update_result.set_result(self.model_update_tmp)
|
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]):
|
async def _resolve_lora_path(self, obj: Union[GenerateReqInput, EmbeddingReqInput]):
|
||||||
if isinstance(obj.lora_path, str):
|
if isinstance(obj.lora_path, str):
|
||||||
unique_lora_paths = set([obj.lora_path])
|
unique_lora_paths = set([obj.lora_path])
|
||||||
@@ -2181,7 +2126,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||||
created_time: Optional[float] = None,
|
created_time: Optional[float] = None,
|
||||||
request: Optional[fastapi.Request] = None,
|
request: Optional[fastapi.Request] = None,
|
||||||
trace_parent: Optional[str] = None,
|
traceparent: Optional[str] = None,
|
||||||
):
|
):
|
||||||
external_trace_header = None
|
external_trace_header = None
|
||||||
if request:
|
if request:
|
||||||
@@ -2189,11 +2134,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
trace_set_remote_propagate_context(request.headers["trace_context"])
|
trace_set_remote_propagate_context(request.headers["trace_context"])
|
||||||
else:
|
else:
|
||||||
external_trace_header = extract_trace_headers(request.headers)
|
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
|
# 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
|
# real request object but we still need to propagate the traceparent from
|
||||||
# the traceparent that is explicitly passed in
|
# the traceparent that is explicitly passed in
|
||||||
external_trace_header = {"trace_parent": trace_parent}
|
external_trace_header = {"traceparent": traceparent}
|
||||||
|
|
||||||
if obj.is_single:
|
if obj.is_single:
|
||||||
bootstrap_room = (
|
bootstrap_room = (
|
||||||
|
|||||||
@@ -309,3 +309,23 @@ class TokenizerManagerMultiItemMixin:
|
|||||||
]
|
]
|
||||||
|
|
||||||
return score_list
|
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
|
||||||
|
|||||||
@@ -5012,12 +5012,12 @@ class PortArgs:
|
|||||||
else:
|
else:
|
||||||
nccl_port = server_args.nccl_port
|
nccl_port = server_args.nccl_port
|
||||||
|
|
||||||
if server_args.tokenizer_worker_num > 1:
|
if server_args.tokenizer_worker_num == 1:
|
||||||
|
tokenizer_worker_ipc_name = None
|
||||||
|
else:
|
||||||
tokenizer_worker_ipc_name = (
|
tokenizer_worker_ipc_name = (
|
||||||
f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
|
f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
tokenizer_worker_ipc_name = None
|
|
||||||
|
|
||||||
if not server_args.enable_dp_attention:
|
if not server_args.enable_dp_attention:
|
||||||
# Normal case, use IPC within a single node
|
# Normal case, use IPC within a single node
|
||||||
|
|||||||
@@ -0,0 +1,116 @@
|
|||||||
|
import glob
|
||||||
|
import os
|
||||||
|
import pickle
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=40, suite="nightly-1-gpu", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCrashDump(CustomTestCase):
|
||||||
|
crash_dump_folder = None
|
||||||
|
MAX_NEW_TOKENS = 4
|
||||||
|
NUM_REQUESTS_BEFORE_CRASH = 5
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.crash_dump_folder = tempfile.mkdtemp(prefix="crash_dump_test_")
|
||||||
|
|
||||||
|
with envs.SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS.override(
|
||||||
|
cls.NUM_REQUESTS_BEFORE_CRASH * cls.MAX_NEW_TOKENS + 10
|
||||||
|
):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
"Qwen/Qwen3-0.6B",
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--crash-dump-folder",
|
||||||
|
cls.crash_dump_folder,
|
||||||
|
"--skip-server-warmup",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def test_crash_dump_generated(self):
|
||||||
|
"""Test that crash dump file is generated after server crash."""
|
||||||
|
# Send multiple requests to trigger the crash
|
||||||
|
for i in range(self.NUM_REQUESTS_BEFORE_CRASH * 2):
|
||||||
|
try:
|
||||||
|
response = requests.post(
|
||||||
|
DEFAULT_URL_FOR_TEST + "/generate",
|
||||||
|
json={
|
||||||
|
"text": f"Hello, this is request {i}.",
|
||||||
|
"sampling_params": {
|
||||||
|
"max_new_tokens": self.MAX_NEW_TOKENS,
|
||||||
|
"temperature": 0,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
except requests.exceptions.RequestException:
|
||||||
|
# Connection error expected after crash
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Wait for crash dump to be written
|
||||||
|
time.sleep(5)
|
||||||
|
|
||||||
|
# Find the crash dump file
|
||||||
|
dump_pattern = os.path.join(self.crash_dump_folder, "*", "crash_dump_*.pkl")
|
||||||
|
dump_files = glob.glob(dump_pattern)
|
||||||
|
|
||||||
|
# Check that a dump file was created
|
||||||
|
self.assertTrue(
|
||||||
|
len(dump_files) > 0,
|
||||||
|
f"No crash dump file found in {self.crash_dump_folder}. "
|
||||||
|
f"Pattern: {dump_pattern}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Read the dump file and verify contents
|
||||||
|
dump_file = dump_files[0]
|
||||||
|
with open(dump_file, "rb") as f:
|
||||||
|
dump_data = pickle.load(f)
|
||||||
|
|
||||||
|
# Verify the dump structure
|
||||||
|
self.assertIn("server_args", dump_data)
|
||||||
|
self.assertIn("requests", dump_data)
|
||||||
|
|
||||||
|
# Check that there are more than 5 requests in the dump
|
||||||
|
requests_list = dump_data["requests"]
|
||||||
|
self.assertGreater(
|
||||||
|
len(requests_list),
|
||||||
|
self.NUM_REQUESTS_BEFORE_CRASH,
|
||||||
|
f"Expected more than {self.NUM_REQUESTS_BEFORE_CRASH} requests in dump, but got {len(requests_list)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify each request tuple has the expected structure (obj, out, created_time, finish_time)
|
||||||
|
for i, req_tuple in enumerate(requests_list):
|
||||||
|
self.assertIsInstance(
|
||||||
|
req_tuple,
|
||||||
|
tuple,
|
||||||
|
f"Request {i} should be a tuple, got {type(req_tuple)}",
|
||||||
|
)
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
len(req_tuple),
|
||||||
|
4,
|
||||||
|
f"Request {i} tuple should have at least 4 elements",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user