Tiny refactor request logger (#15740)

This commit is contained in:
fzyzcjy
2025-12-24 19:51:26 +08:00
committed by GitHub
parent d6108166bc
commit c6a6ba4368
6 changed files with 248 additions and 128 deletions

View File

@@ -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:

View File

@@ -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:

View File

@@ -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],

View File

@@ -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]

View 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)