[diffusion] log: improve logging while multiprocessing (#12997)

This commit is contained in:
Mick
2025-11-12 12:08:37 +08:00
committed by GitHub
parent 151e13687a
commit 60cb716720
4 changed files with 146 additions and 126 deletions
@@ -15,6 +15,8 @@ import requests
import torch import torch
from packaging import version from packaging import version
from sglang.multimodal_gen.runtime.utils.logging_utils import suppress_other_loggers
if version.parse(version.parse(PIL.__version__).base_version) >= version.parse("9.1.0"): if version.parse(version.parse(PIL.__version__).base_version) >= version.parse("9.1.0"):
PIL_INTERPOLATION = { PIL_INTERPOLATION = {
"linear": PIL.Image.Resampling.BILINEAR, "linear": PIL.Image.Resampling.BILINEAR,
@@ -108,7 +110,8 @@ def load_image(
""" """
if isinstance(image, str): if isinstance(image, str):
if image.startswith("http://") or image.startswith("https://"): if image.startswith("http://") or image.startswith("https://"):
image = PIL.Image.open(requests.get(image, stream=True).raw) with suppress_other_loggers(not_suppress_on_main_rank=True):
image = PIL.Image.open(requests.get(image, stream=True).raw)
elif os.path.isfile(image): elif os.path.isfile(image):
image = PIL.Image.open(image) image = PIL.Image.open(image)
else: else:
@@ -9,7 +9,6 @@ from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device, get_local_torch_device,
get_sp_parallel_rank, get_sp_parallel_rank,
get_sp_world_size, get_sp_world_size,
logger,
sequence_model_parallel_all_gather, sequence_model_parallel_all_gather,
) )
from sglang.multimodal_gen.runtime.layers.attention.backends.sliding_tile_attn import ( from sglang.multimodal_gen.runtime.layers.attention.backends.sliding_tile_attn import (
@@ -268,7 +267,7 @@ class DmdDenoisingStage(DenoisingStage):
denoising_loop_end_time = time.time() denoising_loop_end_time = time.time()
if len(timesteps) > 0: if len(timesteps) > 0:
logger.info( self.log_info(
"Average time per step: %.4f seconds", "Average time per step: %.4f seconds",
(denoising_loop_end_time - denoising_loop_start_time) / len(timesteps), (denoising_loop_end_time - denoising_loop_start_time) / len(timesteps),
) )
@@ -34,7 +34,10 @@ from huggingface_hub import snapshot_download
from transformers import AutoConfig, PretrainedConfig from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import (
init_logger,
suppress_other_loggers,
)
logger = init_logger(__name__) logger = init_logger(__name__)
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = { _CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
@@ -370,7 +373,9 @@ def maybe_download_model(
logger.info( logger.info(
"Downloading model snapshot from HF Hub for %s...", model_name_or_path "Downloading model snapshot from HF Hub for %s...", model_name_or_path
) )
with get_lock(model_name_or_path).acquire(poll_interval=2): with get_lock(model_name_or_path).acquire(
poll_interval=2
), suppress_other_loggers(not_suppress_on_main_rank=True):
local_path = snapshot_download( local_path = snapshot_download(
repo_id=model_name_or_path, repo_id=model_name_or_path,
ignore_patterns=["*.onnx", "*.msgpack"], ignore_patterns=["*.onnx", "*.msgpack"],
@@ -5,15 +5,13 @@
"""Logging configuration for sglang.multimodal_gen.""" """Logging configuration for sglang.multimodal_gen."""
import argparse import argparse
import datetime import datetime
import json
import logging import logging
import os import os
import sys import sys
import warnings import warnings
from contextlib import contextmanager
from functools import lru_cache, partial from functools import lru_cache, partial
from logging import Logger from logging import Logger
from logging.config import dictConfig
from os import path
from types import MethodType from types import MethodType
from typing import Any, cast from typing import Any, cast
@@ -29,9 +27,6 @@ GREEN = "\033[92m"
YELLOW = "\033[93m" YELLOW = "\033[93m"
RESET = "\033[0;0m" RESET = "\033[0;0m"
_warned_local_main_process = False
_warned_main_process = False
_FORMAT = ( _FORMAT = (
f"{SGL_DIFFUSION_LOGGING_PREFIX}%(levelname)s %(asctime)s " f"{SGL_DIFFUSION_LOGGING_PREFIX}%(levelname)s %(asctime)s "
"[%(filename)s: %(lineno)d] %(message)s" "[%(filename)s: %(lineno)d] %(message)s"
@@ -129,69 +124,42 @@ def _print_warning_once(logger: Logger, msg: str) -> None:
logger.warning(msg, stacklevel=2) logger.warning(msg, stacklevel=2)
# TODO(will): add env variable to control this process-aware logging behavior def _get_rank_info():
def _info( """Get rank and local rank from environment variables."""
logger: Logger, try:
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
except (KeyError, ValueError):
rank = 0
local_rank = 0
return rank, local_rank
def _log_process_aware(
level: int,
logger_self: Logger,
msg: object, msg: object,
*args: Any, *args: Any,
main_process_only: bool = True, main_process_only: bool,
local_main_process_only: bool = True, local_main_process_only: bool,
**kwargs: Any, **kwargs: Any,
) -> None: ) -> None:
"""Process-aware INFO level logging function. """Helper function to log a message if the process rank matches the criteria."""
rank, local_rank = _get_rank_info()
This function controls logging behavior based on the process rank, allowing for
selective logging from specific processes in a distributed environment.
Args:
logger: The logger instance to use for logging
msg: The message format string to log
*args: Format string arguments
main_process_only: If True, only log if this is the global main process (RANK=0)
local_main_process_only: If True, only log if this is the local main process (LOCAL_RANK=0)
**kwargs: Additional keyword arguments to pass to the logger.log method
- stacklevel: Defaults to 2 to show the original caller's location
Note:
- When both main_process_only and local_main_process_only are True,
the message will be logged only if both conditions are met
- When both are False, the message will be logged from all processes
- By default, only logs from processes with LOCAL_RANK=0
"""
try:
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
except Exception:
local_rank = 0
rank = 0
is_main_process = rank == 0 is_main_process = rank == 0
is_local_main_process = local_rank == 0 is_local_main_process = local_rank == 0
if (main_process_only and is_main_process) or ( should_log = (
local_main_process_only and is_local_main_process not main_process_only
): and not local_main_process_only
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs) or (main_process_only and is_main_process)
or (local_main_process_only and is_local_main_process)
)
global _warned_local_main_process, _warned_main_process if should_log:
# stacklevel=3 to show the original caller's location,
if not _warned_local_main_process and local_main_process_only: # as this function is called by the patched methods.
# logger.warning( logger_self.log(level, msg, *args, stacklevel=3, **kwargs)
# "%sBy default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s",
# GREEN,
# RESET,
# )
_warned_local_main_process = True
if not _warned_main_process and main_process_only and is_main_process:
# logger.warning(
# "%sis_main_process_only is set to True, logging only from the main (RANK==0) process.%s",
# GREEN,
# RESET,
# )
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
class _SGLDiffusionLogger(Logger): class _SGLDiffusionLogger(Logger):
@@ -224,54 +192,34 @@ class _SGLDiffusionLogger(Logger):
main_process_only: bool = True, main_process_only: bool = True,
local_main_process_only: bool = True, local_main_process_only: bool = True,
**kwargs: Any, **kwargs: Any,
) -> None: ) -> None: ...
_info(
self,
msg,
*args,
main_process_only=main_process_only,
local_main_process_only=local_main_process_only,
**kwargs,
)
def debug( # type: ignore[override]
self,
msg: object,
*args: Any,
main_process_only: bool = True,
local_main_process_only: bool = True,
**kwargs: Any,
) -> None: ...
def _configure_sgl_diffusion_root_logger() -> None: def warning( # type: ignore[override]
logging_config = dict[str, Any]() self,
msg: object,
*args: Any,
main_process_only: bool = False,
local_main_process_only: bool = True,
**kwargs: Any,
) -> None: ...
if not SGL_DIFFUSION_CONFIGURE_LOGGING and SGL_DIFFUSION_LOGGING_CONFIG_PATH: def error( # type: ignore[override]
raise RuntimeError( self,
"SGL_DIFFUSION_CONFIGURE_LOGGING evaluated to false, but " msg: object,
"SGL_DIFFUSION_LOGGING_CONFIG_PATH was given. SGL_DIFFUSION_LOGGING_CONFIG_PATH " *args: Any,
"implies SGL_DIFFUSION_CONFIGURE_LOGGING. Please enable " main_process_only: bool = False,
"SGL_DIFFUSION_CONFIGURE_LOGGING or unset SGL_DIFFUSION_LOGGING_CONFIG_PATH." local_main_process_only: bool = True,
) **kwargs: Any,
) -> None: ...
if SGL_DIFFUSION_CONFIGURE_LOGGING:
logging_config = DEFAULT_LOGGING_CONFIG
if SGL_DIFFUSION_LOGGING_CONFIG_PATH:
if not path.exists(SGL_DIFFUSION_LOGGING_CONFIG_PATH):
raise RuntimeError(
"Could not load logging config. File does not exist: %s",
SGL_DIFFUSION_LOGGING_CONFIG_PATH,
)
with open(SGL_DIFFUSION_LOGGING_CONFIG_PATH, encoding="utf-8") as file:
custom_config = json.loads(file.read())
if not isinstance(custom_config, dict):
raise ValueError(
"Invalid logging config. Expected Dict, got %s.",
type(custom_config).__name__,
)
logging_config = custom_config
for formatter in logging_config.get("formatters", {}).values():
# This provides backwards compatibility after #10134.
if formatter.get("class") == "sglang.multimodal_gen.logging.NewLineFormatter":
formatter["class"] = "sglang.multimodal_gen.logging_utils.NewLineFormatter"
if logging_config:
dictConfig(logging_config)
def init_logger(name: str) -> _SGLDiffusionLogger: def init_logger(name: str) -> _SGLDiffusionLogger:
@@ -281,25 +229,59 @@ def init_logger(name: str) -> _SGLDiffusionLogger:
logger = logging.getLogger(name) logger = logging.getLogger(name)
methods_to_patch = { # Patch instance methods
"info_once": _print_info_once, setattr(logger, "info_once", MethodType(_print_info_once, logger))
"warning_once": _print_warning_once, setattr(logger, "warning_once", MethodType(_print_warning_once, logger))
"info": _info,
}
for method_name, method in methods_to_patch.items(): def _create_patched_method(
setattr( level: int,
logger, method_name, MethodType(method, logger) main_process_only_default: bool,
) # type: ignore[arg-type] local_main_process_only_default: bool,
):
def _method(
self: Logger,
msg: object,
*args: Any,
main_process_only: bool = main_process_only_default,
local_main_process_only: bool = local_main_process_only_default,
**kwargs: Any,
) -> None:
_log_process_aware(
level,
self,
msg,
*args,
main_process_only=main_process_only,
local_main_process_only=local_main_process_only,
**kwargs,
)
return _method
setattr(
logger,
"info",
MethodType(_create_patched_method(logging.INFO, True, True), logger),
)
setattr(
logger,
"debug",
MethodType(_create_patched_method(logging.DEBUG, True, True), logger),
)
setattr(
logger,
"warning",
MethodType(_create_patched_method(logging.WARNING, False, True), logger),
)
setattr(
logger,
"error",
MethodType(_create_patched_method(logging.ERROR, False, True), logger),
)
return cast(_SGLDiffusionLogger, logger) return cast(_SGLDiffusionLogger, logger)
# The root logger is initialized when the module is imported.
# This is thread-safe as the module is only imported once,
# guaranteed by the Python GIL.
# _configure_sgl_diffusion_root_logger()
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -395,7 +377,38 @@ def configure_logger(server_args, prefix: str = ""):
set_uvicorn_logging_configs() set_uvicorn_logging_configs()
def suppress_other_loggers(): @contextmanager
def suppress_other_loggers(not_suppress_on_main_rank: bool = False):
"""
A context manager to temporarily suppress specified loggers.
Args:
not_suppress_on_main_rank (bool): If True, loggers will not be
suppressed on the main process (rank 0).
"""
# This is a global setting that we want to apply to all ranks
warnings.filterwarnings( warnings.filterwarnings(
"ignore", category=UserWarning, message="The given NumPy array is not writable" "ignore", category=UserWarning, message="The given NumPy array is not writable"
) )
should_suppress = True
if not_suppress_on_main_rank:
rank, _ = _get_rank_info()
if rank == 0:
should_suppress = False
loggers_to_suppress = ["urllib3"]
original_levels = {}
if should_suppress:
for logger_name in loggers_to_suppress:
logger = logging.getLogger(logger_name)
original_levels[logger_name] = logger.level
logger.setLevel(logging.WARNING)
try:
yield
finally:
if should_suppress:
for logger_name, level in original_levels.items():
logging.getLogger(logger_name).setLevel(level)