[Diffusion] Diffusion model support log-requests (#23049)

Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Thomas
2026-07-05 12:01:26 +03:00
committed by GitHub
co-authored by Claude ronnie_zheng
parent b070cb2ae0
commit addffd7489
5 changed files with 493 additions and 2 deletions
@@ -1,6 +1,6 @@
import pickle
import time
from typing import Any
from typing import Any, Optional
import zmq
import zmq.asyncio
@@ -9,6 +9,9 @@ from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.request_logger import (
DiffusionRequestLogger,
)
logger = init_logger(__name__)
@@ -56,6 +59,7 @@ class SchedulerClient:
self.context = None
self.scheduler_socket = None
self.server_args = None
self.request_logger: Optional[DiffusionRequestLogger] = None
def initialize(self, server_args: ServerArgs):
if self.context is not None and not self.context.closed:
@@ -63,6 +67,7 @@ class SchedulerClient:
self.close()
self.server_args = server_args
self.request_logger = DiffusionRequestLogger.from_server_args(server_args)
self.context = zmq.Context()
self.scheduler_socket = self.context.socket(zmq.REQ)
@@ -80,6 +85,7 @@ class SchedulerClient:
def forward(self, batch: Any, timeout_ms: int | None = None) -> Any:
"""Sends a batch or request to the scheduler and waits for the response."""
self.request_logger.log_received_request(batch)
previous_timeout_ms = None
if timeout_ms is not None:
previous_timeout_ms = self.scheduler_socket.getsockopt(zmq.RCVTIMEO)
@@ -88,6 +94,7 @@ class SchedulerClient:
self.scheduler_socket.send_pyobj(batch)
output_batch = self.scheduler_socket.recv_pyobj()
_materialize_output_batch_file_refs(output_batch)
self.request_logger.log_finished_request(batch, output_batch)
return output_batch
except zmq.error.Again:
logger.error("Timeout waiting for response from scheduler.")
@@ -142,6 +149,7 @@ class AsyncSchedulerClient:
def __init__(self):
self.context = None
self.server_args = None
self.request_logger: Optional[DiffusionRequestLogger] = None
def initialize(self, server_args: ServerArgs):
if self.context is not None and not self.context.closed:
@@ -151,11 +159,13 @@ class AsyncSchedulerClient:
self.close()
self.server_args = server_args
self.request_logger = DiffusionRequestLogger.from_server_args(server_args)
self.context = zmq.asyncio.Context()
logger.debug("AsyncSchedulerClient initialized with zmq.asyncio.Context")
async def forward(self, batch: Any) -> Any:
"""Sends a batch or request to the scheduler and waits for the response."""
self.request_logger.log_received_request(batch)
if self.context is None:
raise RuntimeError(
"AsyncSchedulerClient is not initialized. Call initialize() first."
@@ -175,6 +185,7 @@ class AsyncSchedulerClient:
payload = await socket.recv()
output_batch = pickle.loads(payload)
_materialize_output_batch_file_refs(output_batch)
self.request_logger.log_finished_request(batch, output_batch)
return output_batch
except zmq.error.Again:
logger.error("Timeout waiting for response from scheduler.")
@@ -14,7 +14,7 @@ import sys
import tempfile
from dataclasses import field
from enum import Enum
from typing import Any, Literal, Optional
from typing import Any, List, Literal, Optional
import addict
import yaml
@@ -343,6 +343,10 @@ class ServerArgs(DisaggServerArgsMixin):
# Logging
log_level: str = "info"
log_requests: bool = False
log_requests_level: int = 2
log_requests_format: str = "text"
log_requests_target: Optional[List[str]] = None
uvicorn_access_log_exclude_prefixes: list[str] = field(default_factory=list)
# Tracing
@@ -1652,6 +1656,38 @@ class ServerArgs(DisaggServerArgsMixin):
default=ServerArgs.otlp_traces_endpoint,
help="OTLP collector endpoint when --enable-trace is set. Format: <host>:<port>",
)
parser.add_argument(
"--log-requests",
action="store_true",
help="Log user-facing fields of all requests (default: False). "
"Verbosity is controlled by --log-requests-level.",
)
parser.add_argument(
"--log-requests-level",
type=int,
default=ServerArgs.log_requests_level,
choices=[0, 1, 2, 3],
help="Verbosity level for request logging. "
"0: Log request metadata only (request_id). "
"1: Log metadata + sampling config (seed, steps, guidance, resolution, frames, fps, ...). "
"2: Log metadata + sampling config + prompt/negative prompt (truncated to 2 KiB). "
"3: Log metadata + sampling config + full prompt/negative prompt.",
)
parser.add_argument(
"--log-requests-format",
type=str,
default=ServerArgs.log_requests_format,
choices=["text", "json"],
help="Format for request logging: 'text' (human-readable) or 'json' (structured)",
)
parser.add_argument(
"--log-requests-target",
type=str,
nargs="+",
default=ServerArgs.log_requests_target,
help="Target(s) for request logging: 'stdout' and/or directory path(s) for file output. "
"Can specify multiple targets, e.g., '--log-requests-target stdout /my/path'. ",
)
parser.add_argument(
"--uvicorn-access-log-exclude-prefixes",
type=str,
@@ -0,0 +1,204 @@
# Copyright 2026 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 typing import Any, Optional
from sglang.srt.environ import envs
from sglang.srt.utils.log_utils import create_log_targets, log_json
from sglang.srt.utils.request_logger import (
_dataclass_to_string_truncated,
_transform_data_for_logging,
)
# Core generation knobs logged per record. Prompt text is logged separately,
# gated by the level, so it is excluded here.
_SAMPLING_CONFIG_FIELDS = (
"data_type",
"seed",
"num_inference_steps",
"guidance_scale",
"true_cfg_scale",
"width",
"height",
"num_frames",
"fps",
"num_outputs_per_prompt",
)
# Level 2 truncates prompt text; level 3 keeps it whole. Lower levels log no
# prompt at all.
_TRUNCATE_LENGTH = 2048
_UNLIMITED = 1 << 30
class DiffusionRequestLogger:
def __init__(
self,
log_requests: bool,
log_requests_level: int,
log_requests_format: str,
log_requests_target: Optional[list],
):
self.log_requests = log_requests
self.log_requests_level = log_requests_level
self.log_requests_format = log_requests_format
self.log_requests_target = log_requests_target
self.targets = create_log_targets(
targets=log_requests_target, name_prefix=__name__
)
self.log_exceeded_ms = envs.SGLANG_LOG_REQUEST_EXCEEDED_MS.get()
self._max_length = (
_TRUNCATE_LENGTH if self.log_requests_level == 2 else _UNLIMITED
)
@classmethod
def from_server_args(cls, server_args: Any) -> "DiffusionRequestLogger":
"""Build a logger from server args."""
return cls(
log_requests=server_args.log_requests,
log_requests_level=server_args.log_requests_level,
log_requests_format=server_args.log_requests_format,
log_requests_target=server_args.log_requests_target,
)
@staticmethod
def _request_id(req: Any) -> Optional[str]:
"""The request's id, or ``None`` if absent."""
return getattr(req, "request_id", None)
def _config_view(self, req: Any, *, drop_seed: bool = False) -> dict:
"""Sampling config + (level >= 2) prompt, gated by the log level.
Returns ``{}`` below level 1."""
sp = getattr(req, "sampling_params", None)
if sp is None or self.log_requests_level < 1:
return {}
cfg = {name: getattr(sp, name, None) for name in _SAMPLING_CONFIG_FIELDS}
if drop_seed:
cfg.pop("seed", None)
view: dict = {"sampling_params": cfg}
if self.log_requests_level >= 2:
view["prompt"] = getattr(sp, "prompt", None)
view["negative_prompt"] = getattr(sp, "negative_prompt", None)
return view
def _result_view(self, result: Any) -> dict:
"""Result-side fields for a finished record: latency and error."""
e2e_latency = 0.0
metrics = getattr(result, "metrics", None) if result is not None else None
if metrics is not None:
e2e_latency = getattr(metrics, "total_duration_s", 0.0) or 0.0
return {
"meta_info": {"e2e_latency": e2e_latency},
"error": getattr(result, "error", None) if result is not None else None,
}
def _emit(self, msg: str) -> None:
for target in self.targets:
target.info(msg)
def _per_request_view(self, req: Any) -> dict:
"""Per-output identity within a batch: ``request_id`` plus ``seed`` at
level >= 1."""
sp = getattr(req, "sampling_params", None)
view: dict = {"request_id": self._request_id(req)}
if self.log_requests_level >= 1:
view["seed"] = getattr(sp, "seed", None) if sp is not None else None
return view
def _batch_record(self, reqs: list) -> tuple:
"""Build the ``rid`` / ``obj`` for one forward call"""
rids = [self._request_id(req) or "unknown" for req in reqs]
if len(reqs) == 1:
# Single request: scalar rid + flat dict obj (id + config).
req = reqs[0]
obj = {"request_id": self._request_id(req), **self._config_view(req)}
return rids[0], obj
shared_views = [self._config_view(req, drop_seed=True) for req in reqs]
if all(view == shared_views[0] for view in shared_views):
obj = {
**shared_views[0],
"outputs": [self._per_request_view(req) for req in reqs],
}
else:
# Configs genuinely differ: list each request's full payload verbatim.
obj = [
{"request_id": self._request_id(req), **self._config_view(req)}
for req in reqs
]
return rids, obj
def _loggable(self, req: Any) -> bool:
"""Whether ``req`` should be recorded: logging is on, it's a real
generation request (control messages -- LoRA / weight / stats / shutdown
-- have no ``sampling_params`` and are skipped), and it's not a warmup."""
return (
self.log_requests
and getattr(req, "sampling_params", None) is not None
and not getattr(req, "is_warmup", False)
)
def _logged_reqs(self, batch: Any) -> list:
"""Normalize ``batch`` to a list and drop control / warmup requests."""
reqs = batch if isinstance(batch, (list, tuple)) else [batch]
return [r for r in reqs if self._loggable(r)]
def log_received_request(self, batch: Any) -> None:
reqs = self._logged_reqs(batch)
if not reqs:
return
rid, obj = self._batch_record(reqs)
max_length = self._max_length
if self.log_requests_format == "json":
log_json(
self.targets,
"request.received",
{"rid": rid, "obj": _transform_data_for_logging(obj, max_length)},
)
else:
self._emit(
f"Receive: obj={_dataclass_to_string_truncated(obj, max_length)}"
)
def log_finished_request(self, batch: Any, result: Any) -> None:
reqs = self._logged_reqs(batch)
if not reqs:
return
out = self._result_view(result)
e2e_latency_ms = out["meta_info"]["e2e_latency"] * 1000
if self.log_exceeded_ms > 0 and e2e_latency_ms < self.log_exceeded_ms:
return
rid, obj = self._batch_record(reqs)
max_length = self._max_length
if self.log_requests_format == "json":
log_json(
self.targets,
"request.finished",
{
"rid": rid,
"obj": _transform_data_for_logging(obj, max_length),
"out": _transform_data_for_logging(out, max_length),
},
)
else:
self._emit(
f"Finish: obj={_dataclass_to_string_truncated(obj, max_length)}"
f", out={_dataclass_to_string_truncated(out, max_length)}"
)
@@ -0,0 +1,233 @@
"""
Test request logging for diffusion models.
Tests the --log-requests CLI flags for diffusion model serving,
verifying that request logs are correctly written to stdout and files.
"""
import json
import os
import shutil
import tempfile
import time
from pathlib import Path
import pytest
from openai import OpenAI
from sglang.multimodal_gen.test.server.test_server_utils import ServerManager
from sglang.multimodal_gen.test.test_utils import get_dynamic_server_port
# Test models and prompts
IMAGE_MODEL = "Efficient-Large-Model/Sana_600M_512px_diffusers"
IMAGE_PROMPT = "A beautiful sunset over mountains, oil painting style"
VIDEO_MODEL = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
VIDEO_PROMPT = "A cat playing with a ball"
# Timeout settings
BASE_TIMEOUT = float(os.environ.get("SGLANG_TEST_OPENAI_REQUEST_TIMEOUT_SECS", "600"))
POLL_INTERVAL = 1.0
def _start_server(model: str, log_format: str):
"""Start server with request logging enabled."""
temp_dir = tempfile.mkdtemp()
port = get_dynamic_server_port()
extra_args = (
f"--log-requests "
f"--log-requests-level 2 "
f"--log-requests-format {log_format} "
f"--log-requests-target stdout {temp_dir} "
f"--strict-ports"
)
wait_deadline = float(os.environ.get("SGLANG_TEST_WAIT_SECS", "1200"))
manager = ServerManager(
model=model,
port=port,
wait_deadline=wait_deadline,
extra_args=extra_args,
)
ctx = manager.start()
ctx.temp_dir = temp_dir
return ctx
def _cleanup_server(ctx):
"""Cleanup server and temp directory."""
ctx.cleanup()
shutil.rmtree(ctx.temp_dir, ignore_errors=True)
def _create_client(ctx) -> OpenAI:
"""Create OpenAI client for the server."""
return OpenAI(
api_key="test",
base_url=f"http://localhost:{ctx.port}/v1",
timeout=BASE_TIMEOUT,
)
def _wait_for_video_completion(client: OpenAI, video_id: str, timeout: float):
"""Poll video job until completion."""
deadline = time.time() + timeout
while time.time() < deadline:
page = client.videos.list()
item = next((v for v in page.data if v.id == video_id), None)
status = getattr(item, "status", None) if item else None
if status == "completed":
return True
if status in ("failed", "cancelled", "deleted"):
pytest.fail(f"Video job {video_id} ended with status={status}")
time.sleep(POLL_INTERVAL)
pytest.fail(f"Video job {video_id} did not complete in {timeout}s")
def _verify_json_logs(content: str):
"""Verify JSON logs contain request.received and request.finished events."""
has_received = False
has_finished = False
for line in content.splitlines():
idx = line.find("{")
if idx == -1:
continue
try:
data = json.loads(line[idx:])
except json.JSONDecodeError:
continue
if data.get("event") == "request.received":
has_received = True
elif data.get("event") == "request.finished":
has_finished = True
assert has_received, "request.received event not found"
assert has_finished, "request.finished event not found"
def _verify_text_logs(content: str, prompt: str):
"""Verify text logs contain Receive, prompt, and Finish markers."""
assert "Receive:" in content, "'Receive:' not found"
assert prompt in content, f"Prompt '{prompt}' not found"
assert "Finish:" in content, "'Finish:' not found"
@pytest.fixture(scope="class")
def image_text_server():
"""Server with text-format logging for image model."""
ctx = _start_server(IMAGE_MODEL, "text")
yield ctx
_cleanup_server(ctx)
@pytest.fixture(scope="class")
def image_json_server():
"""Server with JSON-format logging for image model."""
ctx = _start_server(IMAGE_MODEL, "json")
yield ctx
_cleanup_server(ctx)
@pytest.fixture(scope="class")
def video_text_server():
"""Server with text-format logging for video model."""
ctx = _start_server(VIDEO_MODEL, "text")
yield ctx
_cleanup_server(ctx)
@pytest.fixture(scope="class")
def video_json_server():
"""Server with JSON-format logging for video model."""
ctx = _start_server(VIDEO_MODEL, "json")
yield ctx
_cleanup_server(ctx)
class TestImageRequestLoggerText:
"""Test text-format request logging for image models."""
def test_request_logging(self, image_text_server):
ctx = image_text_server
client = _create_client(ctx)
# Image generation is synchronous, waits for completion
client.images.generate(prompt=IMAGE_PROMPT, size="256x256", n=1)
# Verify stdout and file logs
stdout = ctx.log_tail(lines=500)
_verify_text_logs(stdout, IMAGE_PROMPT[:30])
logs = list(Path(ctx.temp_dir).glob("*.log"))
assert logs, "No log files found"
_verify_text_logs("".join(f.read_text() for f in logs), IMAGE_PROMPT[:30])
class TestImageRequestLoggerJson:
"""Test JSON-format request logging for image models."""
def test_request_logging(self, image_json_server):
ctx = image_json_server
client = _create_client(ctx)
# Image generation is synchronous, waits for completion
client.images.generate(prompt=IMAGE_PROMPT, size="256x256", n=1)
# Verify stdout and file logs
stdout = ctx.log_tail(lines=500)
_verify_json_logs(stdout)
logs = list(Path(ctx.temp_dir).glob("*.log"))
assert logs, "No log files found"
_verify_json_logs("".join(f.read_text() for f in logs))
class TestVideoRequestLoggerText:
"""Test text-format request logging for video models."""
def test_request_logging(self, video_text_server):
ctx = video_text_server
client = _create_client(ctx)
# Video generation is async - create job and poll until completion
job = client.videos.create(
prompt=VIDEO_PROMPT,
size="832x480",
extra_body={"num_frames": 5, "num_inference_steps": 10},
)
_wait_for_video_completion(client, job.id, BASE_TIMEOUT * 2)
# Verify stdout and file logs
stdout = ctx.log_tail(lines=500)
_verify_text_logs(stdout, VIDEO_PROMPT[:20])
logs = list(Path(ctx.temp_dir).glob("*.log"))
assert logs, "No log files found"
_verify_text_logs("".join(f.read_text() for f in logs), VIDEO_PROMPT[:20])
class TestVideoRequestLoggerJson:
"""Test JSON-format request logging for video models."""
def test_request_logging(self, video_json_server):
ctx = video_json_server
client = _create_client(ctx)
# Video generation is async - create job and poll until completion
job = client.videos.create(
prompt=VIDEO_PROMPT,
size="832x480",
extra_body={"num_frames": 5, "num_inference_steps": 10},
)
_wait_for_video_completion(client, job.id, BASE_TIMEOUT * 2)
# Verify stdout and file logs
stdout = ctx.log_tail(lines=500)
_verify_json_logs(stdout)
logs = list(Path(ctx.temp_dir).glob("*.log"))
assert logs, "No log files found"
_verify_json_logs("".join(f.read_text() for f in logs))