Tiny refactor request logger (#15740)
This commit is contained in:
@@ -10,11 +10,8 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
|
||||
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
|
||||
from sglang.srt.utils.common import (
|
||||
ceil_align,
|
||||
disable_request_logging,
|
||||
raise_error_or_warn,
|
||||
)
|
||||
from sglang.srt.utils.common import ceil_align, raise_error_or_warn
|
||||
from sglang.srt.utils.request_logger import disable_request_logging
|
||||
from sglang.srt.utils.watchdog import WatchdogRaw
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
@@ -737,44 +737,6 @@ class TokenizerCommunicatorMixin:
|
||||
):
|
||||
await self.send_to_scheduler.send_pyobj(obj)
|
||||
|
||||
def get_log_request_metadata(self):
|
||||
max_length = None
|
||||
skip_names = None
|
||||
out_skip_names = None
|
||||
if self.log_requests:
|
||||
if self.log_requests_level == 0:
|
||||
max_length = 1 << 30
|
||||
skip_names = {
|
||||
"text",
|
||||
"input_ids",
|
||||
"input_embeds",
|
||||
"image_data",
|
||||
"audio_data",
|
||||
"lora_path",
|
||||
"sampling_params",
|
||||
}
|
||||
out_skip_names = {"text", "output_ids", "embedding"}
|
||||
elif self.log_requests_level == 1:
|
||||
max_length = 1 << 30
|
||||
skip_names = {
|
||||
"text",
|
||||
"input_ids",
|
||||
"input_embeds",
|
||||
"image_data",
|
||||
"audio_data",
|
||||
"lora_path",
|
||||
}
|
||||
out_skip_names = {"text", "output_ids", "embedding"}
|
||||
elif self.log_requests_level == 2:
|
||||
max_length = 2048
|
||||
elif self.log_requests_level == 3:
|
||||
max_length = 1 << 30
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid --log-requests-level: {self.log_requests_level=}"
|
||||
)
|
||||
return max_length, skip_names, out_skip_names
|
||||
|
||||
def _update_weight_version_if_provided(self, weight_version: Optional[str]) -> None:
|
||||
"""Update weight version if provided."""
|
||||
if weight_version is not None:
|
||||
|
||||
@@ -98,7 +98,6 @@ from sglang.srt.tracing.trace import (
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
configure_gc_warning,
|
||||
dataclass_to_string_truncated,
|
||||
freeze_gc,
|
||||
get_bool_env_var,
|
||||
get_or_create_event_loop,
|
||||
@@ -111,6 +110,7 @@ from sglang.srt.utils.hf_transformers_utils import (
|
||||
get_tokenizer,
|
||||
get_tokenizer_from_processor,
|
||||
)
|
||||
from sglang.srt.utils.request_logger import RequestLogger
|
||||
from sglang.srt.utils.watchdog import Watchdog
|
||||
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||
|
||||
@@ -182,8 +182,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
# Parse args
|
||||
self.server_args = server_args
|
||||
self.enable_metrics = server_args.enable_metrics
|
||||
self.log_requests = server_args.log_requests
|
||||
self.log_requests_level = server_args.log_requests_level
|
||||
self.request_logger = RequestLogger(
|
||||
log_requests=server_args.log_requests,
|
||||
log_requests_level=server_args.log_requests_level,
|
||||
)
|
||||
self.preferred_sampling_params = server_args.preferred_sampling_params
|
||||
self.crash_dump_folder = server_args.crash_dump_folder
|
||||
self.enable_trace = server_args.enable_trace
|
||||
@@ -308,12 +310,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
self.dump_requests_folder = "" # By default do not dump
|
||||
self.dump_requests_threshold = 1000
|
||||
self.dump_request_list: List[Tuple] = []
|
||||
self.log_request_metadata = self.get_log_request_metadata()
|
||||
self.crash_dump_request_list: deque[Tuple] = deque()
|
||||
self.crash_dump_performed = False # Flag to ensure dump is only called once
|
||||
|
||||
# Initialize performance metrics loggers with proper skip names
|
||||
_, obj_skip_names, out_skip_names = self.log_request_metadata
|
||||
_, obj_skip_names, out_skip_names = self.request_logger.metadata
|
||||
self.request_metrics_exporter_manager = RequestMetricsExporterManager(
|
||||
self.server_args, obj_skip_names, out_skip_names
|
||||
)
|
||||
@@ -432,8 +433,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
self._trace_request_start(obj, created_time, request)
|
||||
if self.server_args.tokenizer_worker_num > 1:
|
||||
self._attach_multi_http_worker_info(obj)
|
||||
if self.log_requests:
|
||||
self._log_received_request(obj)
|
||||
self.request_logger.log_received_request(obj, self.tokenizer)
|
||||
|
||||
async with self.is_pause_cond:
|
||||
await self.is_pause_cond.wait_for(lambda: not self.is_pause)
|
||||
@@ -1044,13 +1044,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
out["meta_info"][
|
||||
"response_sent_to_client_ts"
|
||||
] = state.response_sent_to_client_ts
|
||||
if self.log_requests:
|
||||
max_length, skip_names, out_skip_names = self.log_request_metadata
|
||||
if self.model_config.is_multimodal_gen:
|
||||
msg = f"Finish: obj={dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}"
|
||||
else:
|
||||
msg = f"Finish: obj={dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}, out={dataclass_to_string_truncated(out, max_length, skip_names=out_skip_names)}"
|
||||
logger.info(msg)
|
||||
self.request_logger.log_finished_request(
|
||||
obj, out, is_multimodal_gen=self.model_config.is_multimodal_gen
|
||||
)
|
||||
|
||||
if self.request_metrics_exporter_manager.exporter_enabled():
|
||||
# Asynchronously write metrics for this request using the exporter manager.
|
||||
@@ -1315,10 +1311,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
return all_success, all_message, all_paused_requests
|
||||
|
||||
def configure_logging(self, obj: ConfigureLoggingReq):
|
||||
if obj.log_requests is not None:
|
||||
self.log_requests = obj.log_requests
|
||||
if obj.log_requests_level is not None:
|
||||
self.log_requests_level = obj.log_requests_level
|
||||
self.request_logger.configure(
|
||||
log_requests=obj.log_requests,
|
||||
log_requests_level=obj.log_requests_level,
|
||||
)
|
||||
if obj.dump_requests_folder is not None:
|
||||
self.dump_requests_folder = obj.dump_requests_folder
|
||||
if obj.dump_requests_threshold is not None:
|
||||
@@ -1326,7 +1322,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
if obj.crash_dump_folder is not None:
|
||||
self.crash_dump_folder = obj.crash_dump_folder
|
||||
logging.info(f"Config logging: {obj=}")
|
||||
self.log_request_metadata = self.get_log_request_metadata()
|
||||
|
||||
async def freeze_gc(self):
|
||||
"""Send a freeze_gc message to the scheduler first, then freeze locally."""
|
||||
@@ -2144,23 +2139,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
# Look up the LoRA ID from the registry and start tracking ongoing LoRA requests.
|
||||
obj.lora_id = await self.lora_registry.acquire(obj.lora_path)
|
||||
|
||||
def _log_received_request(self, obj: Union[GenerateReqInput, EmbeddingReqInput]):
|
||||
max_length, skip_names, _ = self.log_request_metadata
|
||||
logger.info(
|
||||
f"Receive: obj={dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}"
|
||||
)
|
||||
|
||||
# FIXME: This is a temporary fix to get the text from the input ids.
|
||||
# We should remove this once we have a proper way.
|
||||
if (
|
||||
self.log_requests_level >= 2
|
||||
and obj.text is None
|
||||
and obj.input_ids is not None
|
||||
and self.tokenizer is not None
|
||||
):
|
||||
decoded = self.tokenizer.decode(obj.input_ids, skip_special_tokens=False)
|
||||
obj.text = decoded
|
||||
|
||||
def _trace_request_start(
|
||||
self,
|
||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||
|
||||
@@ -18,7 +18,6 @@ import argparse
|
||||
import asyncio
|
||||
import builtins
|
||||
import ctypes
|
||||
import dataclasses
|
||||
import functools
|
||||
import importlib
|
||||
import inspect
|
||||
@@ -66,7 +65,6 @@ from typing import (
|
||||
Optional,
|
||||
Protocol,
|
||||
Sequence,
|
||||
Set,
|
||||
Tuple,
|
||||
TypeVar,
|
||||
Union,
|
||||
@@ -2060,53 +2058,6 @@ def set_gpu_proc_affinity(
|
||||
logger.info(f"Process {pid} gpu_id {gpu_id} is running on CPUs: {p.cpu_affinity()}")
|
||||
|
||||
|
||||
@lru_cache(maxsize=2)
|
||||
def disable_request_logging() -> bool:
|
||||
return get_bool_env_var("SGLANG_DISABLE_REQUEST_LOGGING")
|
||||
|
||||
|
||||
def dataclass_to_string_truncated(
|
||||
data, max_length=2048, skip_names: Optional[Set[str]] = None
|
||||
):
|
||||
if skip_names is None:
|
||||
skip_names = set()
|
||||
if isinstance(data, str):
|
||||
if len(data) > max_length:
|
||||
half_length = max_length // 2
|
||||
return f"{repr(data[:half_length])} ... {repr(data[-half_length:])}"
|
||||
else:
|
||||
return f"{repr(data)}"
|
||||
elif isinstance(data, (list, tuple)):
|
||||
if len(data) > max_length:
|
||||
half_length = max_length // 2
|
||||
return str(data[:half_length]) + " ... " + str(data[-half_length:])
|
||||
else:
|
||||
return str(data)
|
||||
elif isinstance(data, dict):
|
||||
return (
|
||||
"{"
|
||||
+ ", ".join(
|
||||
f"'{k}': {dataclass_to_string_truncated(v, max_length)}"
|
||||
for k, v in data.items()
|
||||
if k not in skip_names
|
||||
)
|
||||
+ "}"
|
||||
)
|
||||
elif dataclasses.is_dataclass(data):
|
||||
fields = dataclasses.fields(data)
|
||||
return (
|
||||
f"{data.__class__.__name__}("
|
||||
+ ", ".join(
|
||||
f"{f.name}={dataclass_to_string_truncated(getattr(data, f.name), max_length)}"
|
||||
for f in fields
|
||||
if f.name not in skip_names
|
||||
)
|
||||
+ ")"
|
||||
)
|
||||
else:
|
||||
return str(data)
|
||||
|
||||
|
||||
def permute_weight(x: torch.Tensor) -> torch.Tensor:
|
||||
b_ = x.shape[0]
|
||||
n_ = x.shape[1]
|
||||
|
||||
172
python/sglang/srt/utils/request_logger.py
Normal file
172
python/sglang/srt/utils/request_logger.py
Normal file
@@ -0,0 +1,172 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, Optional, Set, Tuple, Union
|
||||
|
||||
from sglang.srt.utils.common import get_bool_env_var
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RequestLogger:
|
||||
def __init__(self, log_requests: bool, log_requests_level: int):
|
||||
self.log_requests = log_requests
|
||||
self.log_requests_level = log_requests_level
|
||||
self.metadata: Tuple[Optional[int], Optional[Set[str]], Optional[Set[str]]] = (
|
||||
self._compute_metadata()
|
||||
)
|
||||
|
||||
def configure(
|
||||
self,
|
||||
log_requests: Optional[bool] = None,
|
||||
log_requests_level: Optional[int] = None,
|
||||
) -> None:
|
||||
if log_requests is not None:
|
||||
self.log_requests = log_requests
|
||||
if log_requests_level is not None:
|
||||
self.log_requests_level = log_requests_level
|
||||
self.metadata = self._compute_metadata()
|
||||
|
||||
def log_received_request(
|
||||
self, obj: Union["GenerateReqInput", "EmbeddingReqInput"], tokenizer: Any = None
|
||||
) -> None:
|
||||
if not self.log_requests:
|
||||
return
|
||||
|
||||
max_length, skip_names, _ = self.metadata
|
||||
logger.info(
|
||||
f"Receive: obj={_dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}"
|
||||
)
|
||||
|
||||
# FIXME: This is a temporary fix to get the text from the input ids.
|
||||
# We should remove this once we have a proper way.
|
||||
if (
|
||||
self.log_requests_level >= 2
|
||||
and obj.text is None
|
||||
and obj.input_ids is not None
|
||||
and tokenizer is not None
|
||||
):
|
||||
decoded = tokenizer.decode(obj.input_ids, skip_special_tokens=False)
|
||||
obj.text = decoded
|
||||
|
||||
def log_finished_request(
|
||||
self,
|
||||
obj: Union["GenerateReqInput", "EmbeddingReqInput"],
|
||||
out: Any,
|
||||
is_multimodal_gen: bool = False,
|
||||
) -> None:
|
||||
if not self.log_requests:
|
||||
return
|
||||
|
||||
max_length, skip_names, out_skip_names = self.metadata
|
||||
if is_multimodal_gen:
|
||||
msg = f"Finish: obj={_dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}"
|
||||
else:
|
||||
msg = f"Finish: obj={_dataclass_to_string_truncated(obj, max_length, skip_names=skip_names)}, out={_dataclass_to_string_truncated(out, max_length, skip_names=out_skip_names)}"
|
||||
logger.info(msg)
|
||||
|
||||
def _compute_metadata(
|
||||
self,
|
||||
) -> Tuple[Optional[int], Optional[Set[str]], Optional[Set[str]]]:
|
||||
max_length: Optional[int] = None
|
||||
skip_names: Optional[Set[str]] = None
|
||||
out_skip_names: Optional[Set[str]] = None
|
||||
if self.log_requests:
|
||||
if self.log_requests_level == 0:
|
||||
max_length = 1 << 30
|
||||
skip_names = {
|
||||
"text",
|
||||
"input_ids",
|
||||
"input_embeds",
|
||||
"image_data",
|
||||
"audio_data",
|
||||
"lora_path",
|
||||
"sampling_params",
|
||||
}
|
||||
out_skip_names = {"text", "output_ids", "embedding"}
|
||||
elif self.log_requests_level == 1:
|
||||
max_length = 1 << 30
|
||||
skip_names = {
|
||||
"text",
|
||||
"input_ids",
|
||||
"input_embeds",
|
||||
"image_data",
|
||||
"audio_data",
|
||||
"lora_path",
|
||||
}
|
||||
out_skip_names = {"text", "output_ids", "embedding"}
|
||||
elif self.log_requests_level == 2:
|
||||
max_length = 2048
|
||||
elif self.log_requests_level == 3:
|
||||
max_length = 1 << 30
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid --log-requests-level: {self.log_requests_level=}"
|
||||
)
|
||||
return max_length, skip_names, out_skip_names
|
||||
|
||||
|
||||
# TODO remove this?
|
||||
@lru_cache(maxsize=2)
|
||||
def disable_request_logging() -> bool:
|
||||
return get_bool_env_var("SGLANG_DISABLE_REQUEST_LOGGING")
|
||||
|
||||
|
||||
def _dataclass_to_string_truncated(
|
||||
data: Any, max_length: int = 2048, skip_names: Optional[Set[str]] = None
|
||||
) -> str:
|
||||
if skip_names is None:
|
||||
skip_names = set()
|
||||
if isinstance(data, str):
|
||||
if len(data) > max_length:
|
||||
half_length = max_length // 2
|
||||
return f"{repr(data[:half_length])} ... {repr(data[-half_length:])}"
|
||||
else:
|
||||
return f"{repr(data)}"
|
||||
elif isinstance(data, (list, tuple)):
|
||||
if len(data) > max_length:
|
||||
half_length = max_length // 2
|
||||
return str(data[:half_length]) + " ... " + str(data[-half_length:])
|
||||
else:
|
||||
return str(data)
|
||||
elif isinstance(data, dict):
|
||||
return (
|
||||
"{"
|
||||
+ ", ".join(
|
||||
f"'{k}': {_dataclass_to_string_truncated(v, max_length)}"
|
||||
for k, v in data.items()
|
||||
if k not in skip_names
|
||||
)
|
||||
+ "}"
|
||||
)
|
||||
elif dataclasses.is_dataclass(data):
|
||||
fields = dataclasses.fields(data)
|
||||
return (
|
||||
f"{data.__class__.__name__}("
|
||||
+ ", ".join(
|
||||
f"{f.name}={_dataclass_to_string_truncated(getattr(data, f.name), max_length)}"
|
||||
for f in fields
|
||||
if f.name not in skip_names
|
||||
)
|
||||
+ ")"
|
||||
)
|
||||
else:
|
||||
return str(data)
|
||||
Reference in New Issue
Block a user