[diffusion] feat: add metrics support (#19084)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
+2
-1
@@ -1669,7 +1669,8 @@
|
||||
{
|
||||
"group": "References",
|
||||
"pages": [
|
||||
"docs/sglang-diffusion/environment_variables"
|
||||
"docs/sglang-diffusion/environment_variables",
|
||||
"docs/sglang-diffusion/production_metrics"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -7,6 +7,8 @@ SGLang exposes the following metrics via Prometheus. You can enable it by adding
|
||||
|
||||
An example of the monitoring dashboard is available in [examples/monitoring/grafana.json](https://github.com/sgl-project/sglang/blob/main/examples/monitoring/grafana/dashboards/json/sglang-dashboard.json).
|
||||
|
||||
## Language model metrics
|
||||
|
||||
Here is an example of the metrics:
|
||||
|
||||
```text Output
|
||||
@@ -134,6 +136,12 @@ sglang:spec_num_steps{model_name="meta-llama/Llama-3.1-8B-Instruct"} 3.0
|
||||
sglang:spec_num_draft_tokens{model_name="meta-llama/Llama-3.1-8B-Instruct"} 4.0
|
||||
```
|
||||
|
||||
## Diffusion metrics
|
||||
|
||||
SGLang Diffusion exposes request, queue, stage and LoRA metrics with
|
||||
`--enable-metrics`. See [Diffusion production metrics](/docs/sglang-diffusion/production_metrics)
|
||||
for the metric reference, counting semantics and disaggregated scraping setup.
|
||||
|
||||
## Setup Guide
|
||||
|
||||
This section describes how to set up the monitoring stack (Prometheus + Grafana) provided in the `examples/monitoring` directory.
|
||||
|
||||
@@ -110,6 +110,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
|
||||
- `--srt-encoder-timeout {SECONDS}`: Timeout in seconds for HTTP requests to the SGLang encoder server
|
||||
- `--srt-encoder-connection-timeout {SECONDS}`: TCP connection timeout in seconds for SGLang encoder server
|
||||
- `--scheduler-rpc-timeout {SECONDS}`: optional end-to-end deadline for an internal scheduler RPC, including scheduler queue time. It is unset by default so valid long-running and queued video jobs are not failed by the transport layer. Set it only when the deployment requires a bounded request deadline; caller cancellation and server shutdown remain effective without it.
|
||||
- `--enable-metrics`: expose Prometheus metrics at `/metrics` (default: disabled). Includes request counts, queue time, host-side stage timing and LoRA state, separated by role and DP replica. See [Production metrics](/docs/sglang-diffusion/production_metrics) for metric semantics and disaggregated scraping.
|
||||
- `--pe-server-url {HTTPADDRESS}`: url of SGLang server hosting the PE model (e.g., for ERNIE-Image). See [Models with Prompt Enhancement](/docs/sglang-diffusion/models_with_pe).
|
||||
|
||||
### Sampling and output
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
---
|
||||
title: "Production Metrics"
|
||||
description: "Monitor SGLang Diffusion requests, queues, stages and LoRA state with Prometheus."
|
||||
---
|
||||
|
||||
## Enable metrics
|
||||
|
||||
Enable metrics on an NVIDIA CUDA deployment, for example:
|
||||
|
||||
```bash
|
||||
sglang serve --model-path black-forest-labs/FLUX.2-klein-4B \
|
||||
--num-gpus 1 --enable-metrics --port 30000
|
||||
curl http://localhost:30000/metrics
|
||||
```
|
||||
|
||||
Metrics are opt-in. Disabled metrics do not scan queues or collect LoRA status.
|
||||
Enabled metrics add host-side bookkeeping, not GPU synchronization or collectives.
|
||||
|
||||
## Metric reference
|
||||
|
||||
All diffusion metrics carry `role` and `replica` labels. `replica` is the
|
||||
scheduler endpoint; only each DP replica's leader publishes, so TP/SP ranks do
|
||||
not multiply request counts. The table lists additional labels.
|
||||
|
||||
| Metric | Type | Labels | Description |
|
||||
| --- | --- | --- | --- |
|
||||
| `sglang:diffusion_num_queue_reqs` | Gauge | none | Original generation requests waiting for their first dispatch. |
|
||||
| `sglang:diffusion_num_running_reqs` | Gauge | none | Number of diffusion generation requests dispatched by the scheduler and not yet finished. |
|
||||
| `sglang:diffusion_requests_total` | Counter | `status`, `is_warmup` | Completed diffusion generation requests. Dynamic batches are counted per original scheduler request after the merged output is split. |
|
||||
| `sglang:diffusion_request_latency_seconds` | Histogram | `status`, `is_warmup` | Scheduler acceptance to completion, excluding HTTP preprocessing, media encoding and response delivery. |
|
||||
| `sglang:diffusion_queue_time_seconds` | Histogram | `is_warmup` | Time spent waiting in the diffusion scheduler queue. |
|
||||
| `sglang:diffusion_generation_batch_size` | Histogram | `stop_reason` | Generation batch size selected by the diffusion scheduler at dispatch time. |
|
||||
| `sglang:diffusion_stage_host_latency_seconds` | Histogram | `stage` | Host wall time around a stage, not GPU kernel execution time. Step labels are normalized to `DenoisingStep`. |
|
||||
| `sglang:diffusion_lora_loaded_adapters` | Gauge | none | Number of loaded diffusion LoRA adapters. |
|
||||
| `sglang:diffusion_lora_active_modules` | Gauge | none | Number of diffusion modules with active LoRA adapters. |
|
||||
| `sglang:diffusion_lora_active_adapters` | Gauge | none | Number of unique active diffusion LoRA adapters. |
|
||||
| `sglang:diffusion_lora_module_active` | Gauge | `module` | Whether a diffusion module currently has an active LoRA adapter. |
|
||||
|
||||
Request counts refer to original scheduler requests, not generated images,
|
||||
denoising steps or distributed shards. `status` is `success` or `error`;
|
||||
`is_warmup` is `true` or `false`. Queue and running gauges include warmup.
|
||||
Stage observations include warmup and count stage invocations, not requests;
|
||||
asynchronous GPU work can complete in a later stage. For synchronized diagnostic
|
||||
timings, use `SGLANG_DIFFUSION_SYNC_STAGE_PROFILING=1` separately, accepting its
|
||||
synchronization overhead. LoRA gauges update at startup and after LoRA control operations.
|
||||
|
||||
## Disaggregated serving
|
||||
|
||||
In disaggregated serving, the head (`role="server"`) records the original request
|
||||
lifecycle, including role handoffs and errors. Queue time ends at the first
|
||||
encoder dispatch; intermediate waits remain part of request latency. Workers
|
||||
report their own stage and LoRA metrics, not duplicate completed requests.
|
||||
The generation-batch histogram currently describes monolithic scheduling only.
|
||||
|
||||
Single-host pool mode exposes all child metrics through the head's `/metrics`.
|
||||
A scrape aggregates only processes sharing that host's metrics directory.
|
||||
For standalone remote roles, pass `--enable-metrics` to each process and scrape
|
||||
its `--host`/`--port` as well as the head. Each role serves a metrics-only HTTP
|
||||
endpoint. Use a separate, empty `PROMETHEUS_MULTIPROC_DIR` per server launch if
|
||||
you set it yourself; otherwise SGLang creates and owns a temporary directory.
|
||||
Do not share this directory between independent servers or reuse stale files.
|
||||
|
||||
## Query throughput
|
||||
|
||||
Successful, non-warmup request throughput:
|
||||
|
||||
```promql
|
||||
sum(rate(sglang:diffusion_requests_total{status="success",is_warmup="false"}[5m]))
|
||||
```
|
||||
@@ -37,6 +37,7 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import (
|
||||
encode_transfer_msg,
|
||||
is_transfer_message,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import init_metrics
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import (
|
||||
OutputBatch,
|
||||
Req,
|
||||
@@ -173,7 +174,11 @@ class DiffusionServer:
|
||||
self._num_decoders = len(decoder_work_endpoints)
|
||||
self._timeout_s = timeout_s
|
||||
|
||||
self._tracker = RequestTracker()
|
||||
self._tracker = RequestTracker(
|
||||
init_metrics(server_args, role="server")
|
||||
if server_args is not None
|
||||
else None
|
||||
)
|
||||
self._dispatcher = PoolDispatcher(
|
||||
num_encoders=max(1, self._num_encoders),
|
||||
num_denoisers=self._num_denoisers,
|
||||
@@ -476,7 +481,9 @@ class DiffusionServer:
|
||||
request_id = f"ds-{time.monotonic()}"
|
||||
|
||||
try:
|
||||
self._tracker.submit(request_id)
|
||||
self._tracker.submit(
|
||||
request_id, is_warmup=isinstance(req, Req) and req.is_warmup
|
||||
)
|
||||
except ValueError:
|
||||
logger.warning("DiffusionServer: duplicate request_id %s", request_id)
|
||||
return
|
||||
|
||||
@@ -7,6 +7,8 @@ import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import DiffusionMetrics
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -80,16 +82,21 @@ class RequestRecord:
|
||||
class RequestTracker:
|
||||
"""Thread-safe tracker for request state machines."""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, metrics: DiffusionMetrics | None = None):
|
||||
self._metrics = metrics
|
||||
self._lock = threading.Lock()
|
||||
self._requests: dict[str, RequestRecord] = {}
|
||||
|
||||
def submit(self, request_id: str) -> RequestRecord:
|
||||
def submit(self, request_id: str, *, is_warmup: bool = False) -> RequestRecord:
|
||||
with self._lock:
|
||||
if request_id in self._requests:
|
||||
raise ValueError(f"Duplicate request_id: {request_id}")
|
||||
record = RequestRecord(request_id=request_id)
|
||||
self._requests[request_id] = record
|
||||
if self._metrics is not None:
|
||||
self._metrics.enqueue(
|
||||
request_id, is_warmup=is_warmup, now=record.submit_time
|
||||
)
|
||||
return record
|
||||
|
||||
def transition(
|
||||
@@ -124,6 +131,13 @@ class RequestTracker:
|
||||
|
||||
record.state = new_state
|
||||
record.last_transition_time = time.monotonic()
|
||||
if self._metrics is not None:
|
||||
if new_state == RequestState.ENCODER_RUNNING:
|
||||
self._metrics.dispatch(request_id)
|
||||
elif new_state in _TERMINAL_STATES:
|
||||
self._metrics.finish(
|
||||
request_id, error=new_state != RequestState.DONE
|
||||
)
|
||||
if error is not None:
|
||||
record.error = error
|
||||
if encoder_instance is not None:
|
||||
@@ -144,6 +158,8 @@ class RequestTracker:
|
||||
|
||||
def remove(self, request_id: str) -> RequestRecord | None:
|
||||
with self._lock:
|
||||
if self._metrics is not None:
|
||||
self._metrics.finish(request_id, error=True)
|
||||
return self._requests.pop(request_id, None)
|
||||
|
||||
def find_timed_out(self, timeout_s: float) -> list[str]:
|
||||
|
||||
@@ -32,6 +32,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||
prepare_request,
|
||||
save_outputs,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import configure_metrics
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.server_warmup import (
|
||||
@@ -42,6 +43,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
globally_suppress_loggers,
|
||||
init_logger,
|
||||
)
|
||||
from sglang.srt.utils.common import (
|
||||
add_prometheus_middleware,
|
||||
add_prometheus_track_response_middleware,
|
||||
)
|
||||
from sglang.srt.utils.json_response import orjson_response
|
||||
from sglang.version import __version__
|
||||
|
||||
@@ -52,6 +57,7 @@ logger = init_logger(__name__)
|
||||
|
||||
VERTEX_ROUTE = os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate")
|
||||
SERVER_WARMUP_BYPASS_PATHS = (
|
||||
"/metrics",
|
||||
"/liveness",
|
||||
"/health",
|
||||
"/health_generate",
|
||||
@@ -397,6 +403,10 @@ def create_app(server_args: ServerArgs):
|
||||
"""
|
||||
globally_suppress_loggers()
|
||||
app = FastAPI(lifespan=lifespan)
|
||||
if server_args.enable_metrics:
|
||||
configure_metrics()
|
||||
add_prometheus_middleware(app)
|
||||
add_prometheus_track_response_middleware(app)
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
|
||||
@@ -15,6 +15,10 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.entrypoints.control_requests import ShutdownReq
|
||||
from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app
|
||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import (
|
||||
configure_metrics,
|
||||
start_role_metrics_server,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import SchedulerClient
|
||||
from sglang.multimodal_gen.runtime.server_args import (
|
||||
ServerArgs,
|
||||
@@ -137,6 +141,8 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
configure_logger(server_args)
|
||||
|
||||
# Start a new server with multiple worker processes
|
||||
if server_args.enable_metrics:
|
||||
configure_metrics()
|
||||
logger.info("Starting server...")
|
||||
|
||||
# num_gpus is the total world size across every node; each node runs
|
||||
@@ -268,6 +274,8 @@ def launch_pool_disagg_server(
|
||||
configure_logger(server_args)
|
||||
|
||||
num_encoders = len(encoder_gpus)
|
||||
if server_args.enable_metrics:
|
||||
configure_metrics()
|
||||
num_denoisers = len(denoiser_gpus)
|
||||
num_decoders = len(decoder_gpus)
|
||||
logger.info(
|
||||
@@ -423,6 +431,7 @@ def launch_pool_disagg_server(
|
||||
decoder_result_endpoint=decoder_result_ep,
|
||||
dispatch_policy_name=server_args.disagg_dispatch_policy,
|
||||
timeout_s=float(server_args.disagg_timeout),
|
||||
server_args=server_args,
|
||||
)
|
||||
diffusion_server.start()
|
||||
|
||||
@@ -505,6 +514,9 @@ def launch_disagg_server(server_args: ServerArgs):
|
||||
configure_logger(server_args)
|
||||
set_global_server_args(server_args)
|
||||
|
||||
if server_args.enable_metrics:
|
||||
configure_metrics()
|
||||
|
||||
glm_distributed_mode_enabled = (
|
||||
type(server_args.pipeline_config).__name__ == "GlmImagePipelineConfig"
|
||||
and server_args.srt_encoder_url is not None
|
||||
@@ -597,6 +609,8 @@ def launch_disagg_role(server_args: ServerArgs):
|
||||
configure_logger(server_args)
|
||||
|
||||
role_type = server_args.disagg_role
|
||||
if server_args.enable_metrics:
|
||||
configure_metrics()
|
||||
if server_args.disagg_server_addr is None:
|
||||
raise ValueError(
|
||||
f"--disagg-server-addr is required for --disagg-role {role_type.value}"
|
||||
@@ -738,6 +752,8 @@ def launch_disagg_role(server_args: ServerArgs):
|
||||
|
||||
# Block until interrupted
|
||||
try:
|
||||
if server_args.enable_metrics:
|
||||
start_role_metrics_server(server_args)
|
||||
for p in processes:
|
||||
p.join()
|
||||
except KeyboardInterrupt:
|
||||
|
||||
@@ -67,6 +67,10 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.memory_occupation_controller import (
|
||||
MemoryOccupationController,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import (
|
||||
DiffusionMetrics,
|
||||
init_metrics,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import (
|
||||
ComposedPipelineBase,
|
||||
LoRAPipeline,
|
||||
@@ -214,6 +218,8 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
A worker that executes the model on a single GPU.
|
||||
"""
|
||||
|
||||
metrics: DiffusionMetrics | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
local_rank: int,
|
||||
@@ -230,6 +236,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
self.master_port = master_port
|
||||
# FIXME: should we use tcp as distribute init method?
|
||||
self.server_args = server_args
|
||||
self.metrics = init_metrics(server_args, rank)
|
||||
self.pipeline: ComposedPipelineBase = None
|
||||
|
||||
self.init_device_and_model()
|
||||
@@ -260,6 +267,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
# per-rank memory measurements of server warmup forwards; consumed by
|
||||
# the auto-residency placement decision before the server turns ready
|
||||
self._auto_residency_warmup_records: list[WarmupMemoryRecord] = []
|
||||
self._update_lora_metrics()
|
||||
|
||||
def release_realtime_session(self, session_id: str) -> OutputBatch:
|
||||
"""release the session of a realtime connection"""
|
||||
@@ -1458,14 +1466,17 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
"""
|
||||
if not isinstance(self.pipeline, LoRAPipeline):
|
||||
return OutputBatch(error="Lora is not enabled")
|
||||
self.pipeline.set_lora(
|
||||
lora_nickname,
|
||||
lora_path,
|
||||
target,
|
||||
strength,
|
||||
merge_mode=merge_mode,
|
||||
lora_alpha=lora_alpha,
|
||||
)
|
||||
try:
|
||||
self.pipeline.set_lora(
|
||||
lora_nickname,
|
||||
lora_path,
|
||||
target,
|
||||
strength,
|
||||
merge_mode=merge_mode,
|
||||
lora_alpha=lora_alpha,
|
||||
)
|
||||
finally:
|
||||
self._update_lora_metrics()
|
||||
return OutputBatch()
|
||||
|
||||
def merge_lora_weights(
|
||||
@@ -1480,7 +1491,10 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
"""
|
||||
if not isinstance(self.pipeline, LoRAPipeline):
|
||||
return OutputBatch(error="Lora is not enabled")
|
||||
self.pipeline.merge_lora_weights(target, strength)
|
||||
try:
|
||||
self.pipeline.merge_lora_weights(target, strength)
|
||||
finally:
|
||||
self._update_lora_metrics()
|
||||
return OutputBatch()
|
||||
|
||||
def unmerge_lora_weights(self, target: str = "all") -> OutputBatch:
|
||||
@@ -1492,9 +1506,16 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
"""
|
||||
if not isinstance(self.pipeline, LoRAPipeline):
|
||||
return OutputBatch(error="Lora is not enabled")
|
||||
self.pipeline.unmerge_lora_weights(target)
|
||||
try:
|
||||
self.pipeline.unmerge_lora_weights(target)
|
||||
finally:
|
||||
self._update_lora_metrics()
|
||||
return OutputBatch()
|
||||
|
||||
def _update_lora_metrics(self) -> None:
|
||||
if self.metrics is not None and isinstance(self.pipeline, LoRAPipeline):
|
||||
self.metrics.update_lora(self.pipeline.get_lora_status())
|
||||
|
||||
def list_loras(self) -> OutputBatch:
|
||||
"""
|
||||
List loaded LoRA adapters and current application status per module.
|
||||
|
||||
@@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.managers.dynamic_batch_admission import (
|
||||
BatchAdmissionController,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import DiffusionMetrics
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.request_utils import (
|
||||
normalize_output_seeds,
|
||||
@@ -88,6 +89,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
This class does NOT manage worker processes.
|
||||
"""
|
||||
|
||||
metrics: DiffusionMetrics | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
server_args: ServerArgs,
|
||||
@@ -135,6 +138,7 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
server_args=server_args,
|
||||
)
|
||||
self.worker = worker
|
||||
self.metrics = worker.metrics
|
||||
self.gpu_id = gpu_id
|
||||
self._show_warmup_progress = gpu_id == 0
|
||||
self._running = True
|
||||
@@ -658,6 +662,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
reject_reasons: list[str] | None = None,
|
||||
stop_reason: str | None = None,
|
||||
) -> None:
|
||||
if self.metrics is not None:
|
||||
self.metrics.observe_batch(request_count, stop_reason)
|
||||
if not self._batch_metrics_enabled:
|
||||
return
|
||||
|
||||
@@ -793,6 +799,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
output_batch: OutputBatch,
|
||||
) -> None:
|
||||
identity, processed_req = item
|
||||
if self.metrics is not None:
|
||||
self.metrics.finish(id(processed_req), error=output_batch.error is not None)
|
||||
is_warmup = is_warmup_req(processed_req)
|
||||
self._log_warmup_result(output_batch, processed_req, is_warmup)
|
||||
|
||||
@@ -1255,6 +1263,13 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
self.waiting_queue.extend(
|
||||
[(identity, req, now) for identity, req in new_reqs]
|
||||
)
|
||||
if self.metrics is not None:
|
||||
for _, req_or_group in new_reqs:
|
||||
req = get_first_generation_req(req_or_group)
|
||||
if req is not None:
|
||||
self.metrics.enqueue(
|
||||
id(req_or_group), is_warmup=req.is_warmup, now=now
|
||||
)
|
||||
# Reset error count on success
|
||||
self._consecutive_error_count = 0
|
||||
except Exception as e:
|
||||
@@ -1288,6 +1303,14 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
time.sleep(remaining_ms / 1000.0)
|
||||
continue
|
||||
|
||||
if self.metrics is not None:
|
||||
for _, req in items:
|
||||
self.metrics.dispatch(id(req))
|
||||
if (
|
||||
isinstance(req, list)
|
||||
and get_first_generation_req(req) is not None
|
||||
):
|
||||
self.metrics.observe_batch(1, "request_group")
|
||||
try:
|
||||
with maybe_record_function(
|
||||
f"REQ {self._req_label(items)} dispatch+forward"
|
||||
@@ -1305,6 +1328,10 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
self._return_results_sequentially(items, handler_result.outputs)
|
||||
except zmq.ZMQError as e:
|
||||
logger.error(f"ZMQ error sending replies sequentially: {e}")
|
||||
finally:
|
||||
if self.metrics is not None:
|
||||
for _, req in items:
|
||||
self.metrics.finish(id(req), error=True)
|
||||
continue
|
||||
|
||||
if isinstance(handler_result, list):
|
||||
@@ -1336,6 +1363,10 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
# Reply failed; log and keep loop alive to accept future requests
|
||||
logger.error(f"ZMQ error sending reply: {e}")
|
||||
continue
|
||||
finally:
|
||||
if self.metrics is not None:
|
||||
for _, req in items:
|
||||
self.metrics.finish(id(req), error=True)
|
||||
|
||||
self._log_batch_metrics_summary()
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,238 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Opt-in, process-local diffusion metrics; only replica leaders publish."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Hashable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prometheus_client import CollectorRegistry
|
||||
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
_multiproc_dir: tempfile.TemporaryDirectory | None = None
|
||||
_metrics: DiffusionMetrics | None = None
|
||||
|
||||
|
||||
def configure_metrics() -> None:
|
||||
"""Set the multiprocess directory before importing Prometheus or spawning."""
|
||||
global _multiproc_dir
|
||||
if "PROMETHEUS_MULTIPROC_DIR" not in os.environ:
|
||||
_multiproc_dir = tempfile.TemporaryDirectory(prefix="sglang-diffusion-metrics-")
|
||||
os.environ["PROMETHEUS_MULTIPROC_DIR"] = _multiproc_dir.name
|
||||
|
||||
|
||||
def init_metrics(
|
||||
server_args: ServerArgs, rank: int = 0, *, role: str | None = None
|
||||
) -> DiffusionMetrics | None:
|
||||
global _metrics
|
||||
group_size = max(1, server_args.num_gpus // (server_args.dp_size or 1))
|
||||
if not server_args.enable_metrics or rank % group_size:
|
||||
_metrics = None
|
||||
return None
|
||||
replica = rank // group_size
|
||||
_metrics = DiffusionMetrics(
|
||||
role=role or server_args.disagg_role.value,
|
||||
replica=server_args.scheduler_endpoint_for(replica),
|
||||
)
|
||||
return _metrics
|
||||
|
||||
|
||||
def get_metrics() -> DiffusionMetrics | None:
|
||||
return _metrics
|
||||
|
||||
|
||||
def start_role_metrics_server(server_args: ServerArgs) -> None:
|
||||
# prometheus selects its storage backend at import time, after configure_metrics
|
||||
from prometheus_client import CollectorRegistry, multiprocess, start_http_server
|
||||
|
||||
registry = CollectorRegistry()
|
||||
multiprocess.MultiProcessCollector(registry)
|
||||
start_http_server(server_args.port, addr=server_args.host, registry=registry)
|
||||
|
||||
|
||||
class DiffusionMetrics:
|
||||
"""Callbacks are serialized by the scheduler or RequestTracker lock.
|
||||
|
||||
Keys identify original requests, not generated outputs or GPU batches.
|
||||
Request IDs, prompts and adapter paths are never exported as labels.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, *, role: str, replica: str, registry: CollectorRegistry | None = None
|
||||
):
|
||||
# defer this import until the launcher has configured multiprocess storage
|
||||
from prometheus_client import Counter, Gauge, Histogram
|
||||
|
||||
labels = ("role", "replica")
|
||||
self._labels = (role, replica)
|
||||
self._requests: dict[Hashable, tuple[float, bool, bool]] = {}
|
||||
self._queued = 0
|
||||
self._running = 0
|
||||
self._observed_modules: set[str] = set()
|
||||
self.queue = Gauge(
|
||||
"sglang:diffusion_num_queue_reqs",
|
||||
"Requests waiting for first dispatch.",
|
||||
labels,
|
||||
multiprocess_mode="mostrecent",
|
||||
registry=registry,
|
||||
)
|
||||
self.running = Gauge(
|
||||
"sglang:diffusion_num_running_reqs",
|
||||
"Dispatched requests not yet completed.",
|
||||
labels,
|
||||
multiprocess_mode="mostrecent",
|
||||
registry=registry,
|
||||
)
|
||||
self.requests = Counter(
|
||||
"sglang:diffusion_requests_total",
|
||||
"Completed original client requests.",
|
||||
labels + ("status", "is_warmup"),
|
||||
registry=registry,
|
||||
)
|
||||
buckets = (0.01, 0.05, 0.1, 0.5, 1, 2, 5, 10, 20, 30, 60, 120, 300, 600, 1200)
|
||||
self.latency = Histogram(
|
||||
"sglang:diffusion_request_latency_seconds",
|
||||
"Time from scheduler acceptance to completion, excluding HTTP postprocessing.",
|
||||
labels + ("status", "is_warmup"),
|
||||
buckets=buckets,
|
||||
registry=registry,
|
||||
)
|
||||
self.queue_time = Histogram(
|
||||
"sglang:diffusion_queue_time_seconds",
|
||||
"Time until first dispatch.",
|
||||
labels + ("is_warmup",),
|
||||
buckets=buckets,
|
||||
registry=registry,
|
||||
)
|
||||
self.batch_size = Histogram(
|
||||
"sglang:diffusion_generation_batch_size",
|
||||
"Original requests per dispatched batch.",
|
||||
labels + ("stop_reason",),
|
||||
buckets=(1, 2, 4, 8, 16, 32, 64),
|
||||
registry=registry,
|
||||
)
|
||||
self.stage_latency = Histogram(
|
||||
"sglang:diffusion_stage_host_latency_seconds",
|
||||
"Host wall time around a pipeline stage; not synchronized GPU execution time.",
|
||||
labels + ("stage",),
|
||||
buckets=(
|
||||
0.001,
|
||||
0.005,
|
||||
0.01,
|
||||
0.05,
|
||||
0.1,
|
||||
0.5,
|
||||
1,
|
||||
2,
|
||||
5,
|
||||
10,
|
||||
30,
|
||||
60,
|
||||
120,
|
||||
300,
|
||||
1200,
|
||||
),
|
||||
registry=registry,
|
||||
)
|
||||
self.lora_loaded = Gauge(
|
||||
"sglang:diffusion_lora_loaded_adapters",
|
||||
"Loaded LoRA adapters.",
|
||||
labels,
|
||||
multiprocess_mode="mostrecent",
|
||||
registry=registry,
|
||||
)
|
||||
self.lora_modules = Gauge(
|
||||
"sglang:diffusion_lora_active_modules",
|
||||
"Modules with active LoRA adapters.",
|
||||
labels,
|
||||
multiprocess_mode="mostrecent",
|
||||
registry=registry,
|
||||
)
|
||||
self.lora_adapters = Gauge(
|
||||
"sglang:diffusion_lora_active_adapters",
|
||||
"Unique active LoRA adapters.",
|
||||
labels,
|
||||
multiprocess_mode="mostrecent",
|
||||
registry=registry,
|
||||
)
|
||||
self.lora_module = Gauge(
|
||||
"sglang:diffusion_lora_module_active",
|
||||
"Whether a module has an active adapter.",
|
||||
labels + ("module",),
|
||||
multiprocess_mode="mostrecent",
|
||||
registry=registry,
|
||||
)
|
||||
self._publish_depths()
|
||||
|
||||
def _publish_depths(self):
|
||||
self.queue.labels(*self._labels).set(self._queued)
|
||||
self.running.labels(*self._labels).set(self._running)
|
||||
|
||||
def enqueue(self, key: Hashable, *, is_warmup: bool, now: float | None = None):
|
||||
self._requests[key] = (
|
||||
time.monotonic() if now is None else now,
|
||||
is_warmup,
|
||||
False,
|
||||
)
|
||||
self._queued += 1
|
||||
self._publish_depths()
|
||||
|
||||
def dispatch(self, key: Hashable):
|
||||
state = self._requests.get(key)
|
||||
if state is None or state[2]:
|
||||
return
|
||||
start, is_warmup, _ = state
|
||||
self._requests[key] = (start, is_warmup, True)
|
||||
self._queued -= 1
|
||||
self._running += 1
|
||||
self.queue_time.labels(*self._labels, str(is_warmup).lower()).observe(
|
||||
max(0.0, time.monotonic() - start)
|
||||
)
|
||||
self._publish_depths()
|
||||
|
||||
def finish(self, key: Hashable, *, error: bool):
|
||||
state = self._requests.pop(key, None)
|
||||
if state is None:
|
||||
return
|
||||
start, is_warmup, dispatched = state
|
||||
if dispatched:
|
||||
self._running -= 1
|
||||
else:
|
||||
self._queued -= 1
|
||||
labels = (
|
||||
*self._labels,
|
||||
"error" if error else "success",
|
||||
str(is_warmup).lower(),
|
||||
)
|
||||
self.requests.labels(*labels).inc()
|
||||
self.latency.labels(*labels).observe(max(0.0, time.monotonic() - start))
|
||||
self._publish_depths()
|
||||
|
||||
def observe_batch(self, size: int, stop_reason: str | None):
|
||||
reason = (stop_reason or "unspecified").partition(":")[0]
|
||||
self.batch_size.labels(*self._labels, reason).observe(size)
|
||||
|
||||
def observe_stage(self, name: str, seconds: float):
|
||||
if name.startswith("denoising_step_"):
|
||||
name = "DenoisingStep"
|
||||
self.stage_latency.labels(*self._labels, name).observe(seconds)
|
||||
|
||||
def update_lora(self, status: dict):
|
||||
active = status["active"]
|
||||
adapters = {
|
||||
nickname
|
||||
for entries in active.values()
|
||||
for entry in entries
|
||||
for nickname in entry["nicknames"]
|
||||
if nickname
|
||||
}
|
||||
self.lora_loaded.labels(*self._labels).set(len(status["loaded_adapters"]))
|
||||
self.lora_modules.labels(*self._labels).set(len(active))
|
||||
self.lora_adapters.labels(*self._labels).set(len(adapters))
|
||||
self._observed_modules.update(active)
|
||||
for module in self._observed_modules:
|
||||
self.lora_module.labels(*self._labels, module).set(int(module in active))
|
||||
@@ -1403,9 +1403,17 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
if not self._is_lora_effective_for_module(module_name, lora_layers):
|
||||
return None
|
||||
else:
|
||||
nicknames, strengths = self.cur_adapter_config.get(
|
||||
module_name,
|
||||
(
|
||||
[self.cur_adapter_name.get(module_name, None)],
|
||||
[self.cur_adapter_strength.get(module_name, None)],
|
||||
),
|
||||
)
|
||||
return [
|
||||
{
|
||||
"nickname": self.cur_adapter_name.get(module_name, None),
|
||||
"nicknames": nicknames,
|
||||
"path": self.cur_adapter_path.get(module_name, None),
|
||||
"merged": self.is_lora_merged.get(module_name, False),
|
||||
"mode": (
|
||||
@@ -1414,6 +1422,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
else "unmerged"
|
||||
),
|
||||
"strength": self.cur_adapter_strength.get(module_name, None),
|
||||
"strengths": strengths,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@@ -507,6 +507,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
# http server endpoint config
|
||||
host: str | None = "127.0.0.1"
|
||||
port: int | None = 30000
|
||||
enable_metrics: bool = False
|
||||
|
||||
# TODO: webui and their endpoint, check if webui_port is available.
|
||||
webui: bool = False
|
||||
@@ -1366,9 +1367,8 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
)
|
||||
|
||||
def _adjust_network_ports(self):
|
||||
# Disagg role instances (encoder/denoiser/decoder) don't serve HTTP,
|
||||
# so skip settling the HTTP port to avoid unnecessary port collisions.
|
||||
needs_http = self.disagg_role in (
|
||||
# standalone roles only need an HTTP port when exposing metrics
|
||||
needs_http = self.enable_metrics or self.disagg_role in (
|
||||
RoleType.MONOLITHIC,
|
||||
RoleType.SERVER,
|
||||
)
|
||||
@@ -2793,6 +2793,12 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
default=ServerArgs.port,
|
||||
help="Port for the HTTP API server.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable-metrics",
|
||||
action=StoreBoolean,
|
||||
default=ServerArgs.enable_metrics,
|
||||
help="Expose Prometheus metrics at /metrics.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--strict-ports",
|
||||
action=StoreBoolean,
|
||||
|
||||
@@ -17,6 +17,7 @@ from dateutil.tz import UTC
|
||||
|
||||
import sglang
|
||||
import sglang.multimodal_gen.envs as envs
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import get_metrics
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
CYAN,
|
||||
@@ -294,6 +295,7 @@ class StageProfiler:
|
||||
record_as_step: bool = False,
|
||||
):
|
||||
self.stage_name = stage_name
|
||||
self.prometheus = get_metrics()
|
||||
self.metrics = metrics
|
||||
self.logger = logger
|
||||
self.start_time = 0.0
|
||||
@@ -335,14 +337,22 @@ class StageProfiler:
|
||||
msg += f" ({round(available_memory, 2)} GB left)"
|
||||
self.logger.info(msg)
|
||||
|
||||
if (self.log_timing and self.metrics) or self.log_stage_start_end:
|
||||
if (
|
||||
(self.log_timing and self.metrics)
|
||||
or self.log_stage_start_end
|
||||
or self.prometheus is not None
|
||||
):
|
||||
self._maybe_sync_device()
|
||||
self.start_time = time.perf_counter()
|
||||
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
if not ((self.log_timing and self.metrics) or self.log_stage_start_end):
|
||||
if not (
|
||||
(self.log_timing and self.metrics)
|
||||
or self.log_stage_start_end
|
||||
or self.prometheus is not None
|
||||
):
|
||||
return False
|
||||
|
||||
self._maybe_sync_device()
|
||||
@@ -363,6 +373,9 @@ class StageProfiler:
|
||||
f"[{self.stage_name}] finished in {execution_time_s:.4f} seconds",
|
||||
)
|
||||
|
||||
if self.prometheus is not None:
|
||||
self.prometheus.observe_stage(self.stage_name, execution_time_s)
|
||||
|
||||
if self.log_timing and self.metrics:
|
||||
if self._should_record_as_step():
|
||||
self.metrics.record_step(execution_time_s)
|
||||
|
||||
@@ -0,0 +1,286 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Prometheus lifecycle, replica ownership and disabled-path regressions."""
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from collections import deque
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from prometheus_client import CollectorRegistry
|
||||
|
||||
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||
from sglang.multimodal_gen.runtime.disaggregation.request_state import (
|
||||
RequestState,
|
||||
RequestTracker,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
||||
from sglang.multimodal_gen.runtime.managers.scheduler import (
|
||||
Scheduler,
|
||||
_SequentiallyReturnedOutputs,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.observability import metrics as metrics_module
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import DiffusionMetrics
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils import perf_logger
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def metrics():
|
||||
registry = CollectorRegistry()
|
||||
collector = DiffusionMetrics(role="monolithic", replica="0", registry=registry)
|
||||
return collector, registry
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", [False, True])
|
||||
def test_role_checks_http_port_only_when_exporting_metrics(enabled, monkeypatch):
|
||||
args = ServerArgs.__new__(ServerArgs)
|
||||
args.disagg_role = RoleType.DENOISER
|
||||
args.enable_metrics = enabled
|
||||
args.strict_ports = True
|
||||
require_port = Mock()
|
||||
monkeypatch.setattr(args, "_require_port", require_port)
|
||||
args._adjust_network_ports()
|
||||
http_checks = [
|
||||
call.args for call in require_port.call_args_list if call.args[1] == "HTTP"
|
||||
]
|
||||
assert http_checks == ([(args.port, "HTTP")] if enabled else [])
|
||||
|
||||
|
||||
def sample(registry, name, **labels):
|
||||
return registry.get_sample_value(
|
||||
"sglang:diffusion_" + name, {"role": "monolithic", "replica": "0", **labels}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sequential", [False, True])
|
||||
@pytest.mark.parametrize("failure", [False, True])
|
||||
@pytest.mark.parametrize("enabled", [False, True])
|
||||
def test_scheduler_counts_original_requests_and_cleans_up(
|
||||
metrics, sequential, failure, enabled
|
||||
):
|
||||
collector, registry = metrics
|
||||
scheduler = Scheduler.__new__(Scheduler)
|
||||
scheduler.metrics = collector if enabled else None
|
||||
scheduler._disagg_role = RoleType.MONOLITHIC
|
||||
scheduler._disagg_metrics = None
|
||||
scheduler.receiver = None
|
||||
scheduler.context = Mock()
|
||||
|
||||
class NoScanQueue(deque):
|
||||
def __iter__(self):
|
||||
raise AssertionError("metrics must not scan the waiting queue")
|
||||
|
||||
scheduler.waiting_queue = NoScanQueue()
|
||||
scheduler._running = True
|
||||
scheduler._consecutive_error_count = 0
|
||||
scheduler._max_consecutive_errors = 1
|
||||
scheduler._log_warmup_result = Mock()
|
||||
scheduler._log_batch_metrics_summary = Mock()
|
||||
scheduler._cleanup_disagg = Mock()
|
||||
scheduler.return_result = Mock()
|
||||
scheduler.process_received_reqs_with_req_based_warmup = lambda reqs: reqs
|
||||
reqs = [Req(sampling_params=SamplingParams(prompt="test")) for _ in range(2)]
|
||||
# one multi-output request is still one original scheduler request
|
||||
group = [Req(sampling_params=SamplingParams(prompt="group")) for _ in range(3)]
|
||||
reqs.append(group)
|
||||
scheduler.recv_reqs = lambda: [(None, req) for req in reqs]
|
||||
scheduler.get_next_batch_to_run = lambda: [(None, req) for req in reqs]
|
||||
|
||||
def dispatch(items):
|
||||
if enabled:
|
||||
assert sample(registry, "num_running_reqs") == 3
|
||||
assert sample(registry, "num_queue_reqs") == 0
|
||||
scheduler._running = False
|
||||
if failure and not sequential:
|
||||
raise RuntimeError("forward failed")
|
||||
|
||||
def outputs():
|
||||
for index in range(len(items)):
|
||||
if failure and index == 1:
|
||||
raise RuntimeError("forward failed")
|
||||
yield OutputBatch()
|
||||
|
||||
return (
|
||||
_SequentiallyReturnedOutputs(outputs()) if sequential else list(outputs())
|
||||
)
|
||||
|
||||
scheduler._dispatch_items = dispatch
|
||||
scheduler.event_loop()
|
||||
if not enabled:
|
||||
assert not collector._requests
|
||||
assert (
|
||||
sample(registry, "requests_total", status="success", is_warmup="false")
|
||||
is None
|
||||
)
|
||||
return
|
||||
errors = 2 if sequential else 3
|
||||
assert sample(registry, "requests_total", status="success", is_warmup="false") == (
|
||||
(1 if sequential else None) if failure else 3
|
||||
)
|
||||
if failure:
|
||||
assert (
|
||||
sample(registry, "requests_total", status="error", is_warmup="false")
|
||||
== errors
|
||||
)
|
||||
assert sample(registry, "num_running_reqs") == 0
|
||||
assert sample(registry, "num_queue_reqs") == 0
|
||||
assert not collector._requests
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"terminal", [RequestState.DONE, RequestState.FAILED, RequestState.TIMED_OUT]
|
||||
)
|
||||
def test_disagg_lifecycle_counts_once_including_retries(metrics, terminal):
|
||||
collector, registry = metrics
|
||||
tracker = RequestTracker(collector)
|
||||
tracker.submit("request", is_warmup=True)
|
||||
for state in (
|
||||
RequestState.ENCODER_RUNNING,
|
||||
RequestState.ENCODER_DONE,
|
||||
RequestState.DENOISING_RUNNING,
|
||||
RequestState.DENOISING_WAITING,
|
||||
RequestState.DENOISING_RUNNING,
|
||||
RequestState.DENOISING_DONE,
|
||||
RequestState.DECODER_RUNNING,
|
||||
terminal,
|
||||
):
|
||||
tracker.transition("request", state)
|
||||
tracker.remove("request")
|
||||
assert (
|
||||
sample(
|
||||
registry,
|
||||
"requests_total",
|
||||
status="success" if terminal == RequestState.DONE else "error",
|
||||
is_warmup="true",
|
||||
)
|
||||
== 1
|
||||
)
|
||||
assert sample(registry, "queue_time_seconds_count", is_warmup="true") == 1
|
||||
assert sample(registry, "num_running_reqs") == 0
|
||||
tracker.submit("cancelled")
|
||||
tracker.remove("cancelled")
|
||||
assert sample(registry, "num_queue_reqs") == 0
|
||||
assert sample(registry, "requests_total", status="error", is_warmup="false") == 1
|
||||
|
||||
|
||||
def test_only_replica_leaders_construct_metrics(monkeypatch):
|
||||
construct = Mock()
|
||||
monkeypatch.setattr(metrics_module, "DiffusionMetrics", construct)
|
||||
monkeypatch.setattr(metrics_module, "_metrics", None)
|
||||
args = SimpleNamespace(
|
||||
num_gpus=4,
|
||||
dp_size=2,
|
||||
enable_metrics=False,
|
||||
disagg_role=RoleType.MONOLITHIC,
|
||||
scheduler_endpoint_for=lambda replica: f"tcp://localhost:{5555 + replica}",
|
||||
)
|
||||
for rank in range(4):
|
||||
assert metrics_module.init_metrics(args, rank) is None
|
||||
construct.assert_not_called()
|
||||
args.enable_metrics = True
|
||||
for rank in range(4):
|
||||
metrics_module.init_metrics(args, rank)
|
||||
assert construct.call_count == 2
|
||||
assert [call.kwargs["replica"] for call in construct.call_args_list] == [
|
||||
"tcp://localhost:5555",
|
||||
"tcp://localhost:5556",
|
||||
]
|
||||
|
||||
|
||||
def test_disabled_metrics_skip_status_collection_and_timing(monkeypatch):
|
||||
worker = GPUWorker.__new__(GPUWorker)
|
||||
worker.metrics = None
|
||||
worker.pipeline = Mock()
|
||||
worker._update_lora_metrics()
|
||||
worker.pipeline.get_lora_status.assert_not_called()
|
||||
monkeypatch.setattr(metrics_module, "_metrics", None)
|
||||
monkeypatch.setenv("SGLANG_DIFFUSION_STAGE_LOGGING", "0")
|
||||
timer = Mock(side_effect=AssertionError("disabled metrics must not time stages"))
|
||||
monkeypatch.setattr(perf_logger.time, "perf_counter", timer)
|
||||
with perf_logger.StageProfiler("test", Mock(), None):
|
||||
pass
|
||||
timer.assert_not_called()
|
||||
|
||||
|
||||
def test_stage_metrics_do_not_synchronize_and_bound_step_labels(metrics, monkeypatch):
|
||||
collector, registry = metrics
|
||||
monkeypatch.setattr(metrics_module, "_metrics", collector)
|
||||
monkeypatch.delenv("SGLANG_DIFFUSION_SYNC_STAGE_PROFILING", raising=False)
|
||||
device = Mock()
|
||||
monkeypatch.setattr(perf_logger.torch, "get_device_module", lambda: device)
|
||||
for step in range(3):
|
||||
with perf_logger.StageProfiler(f"denoising_step_{step}", Mock(), None):
|
||||
pass
|
||||
device.synchronize.assert_not_called()
|
||||
assert (
|
||||
sample(registry, "stage_host_latency_seconds_count", stage="DenoisingStep") == 3
|
||||
)
|
||||
|
||||
|
||||
def test_lora_deduplicates_adapters_and_resets_inactive_modules(metrics):
|
||||
collector, registry = metrics
|
||||
active = {
|
||||
"transformer": [{"nicknames": ["a", "b"]}],
|
||||
"transformer_2": [{"nicknames": ["a"]}],
|
||||
}
|
||||
collector.update_lora({"loaded_adapters": ["a", "b"], "active": active})
|
||||
assert sample(registry, "lora_active_adapters") == 2
|
||||
collector.update_lora({"loaded_adapters": ["a", "b"], "active": {}})
|
||||
assert sample(registry, "lora_active_modules") == 0
|
||||
assert sample(registry, "lora_module_active", module="transformer") == 0
|
||||
|
||||
|
||||
def test_multiprocess_scrape_keeps_role_and_replica_gauges_separate(tmp_path):
|
||||
env = {**os.environ, "PROMETHEUS_MULTIPROC_DIR": str(tmp_path)}
|
||||
worker = """
|
||||
import sys
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import DiffusionMetrics
|
||||
m = DiffusionMetrics(role=sys.argv[1], replica=sys.argv[2])
|
||||
for i in range(int(sys.argv[3])):
|
||||
m.enqueue(i, is_warmup=False)
|
||||
m.dispatch(0)
|
||||
m.finish(0, error=False)
|
||||
"""
|
||||
for role, replica, count in [
|
||||
("monolithic", "0", 2),
|
||||
("monolithic", "1", 3),
|
||||
("decoder", "0", 1),
|
||||
]:
|
||||
subprocess.run(
|
||||
[sys.executable, "-c", worker, role, replica, str(count)],
|
||||
env=env,
|
||||
check=True,
|
||||
timeout=90,
|
||||
)
|
||||
scrape = """
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from fastapi.testclient import TestClient
|
||||
from prometheus_client import CollectorRegistry, generate_latest, multiprocess
|
||||
from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app
|
||||
r = CollectorRegistry()
|
||||
multiprocess.MultiProcessCollector(r)
|
||||
assert r.get_sample_value('sglang:diffusion_num_queue_reqs', {'role':'monolithic','replica':'0'}) == 1
|
||||
assert r.get_sample_value('sglang:diffusion_num_queue_reqs', {'role':'monolithic','replica':'1'}) == 2
|
||||
assert r.get_sample_value('sglang:diffusion_num_queue_reqs', {'role':'decoder','replica':'0'}) == 0
|
||||
assert b'sglang:diffusion_requests_total' in generate_latest(r)
|
||||
args = SimpleNamespace(enable_metrics=True, pipeline_config=SimpleNamespace(
|
||||
supports_action_endpoint=lambda: False, supports_openpi_endpoint=lambda: False))
|
||||
app = create_app(args)
|
||||
app.state.server_warmup_done = asyncio.Event()
|
||||
response = TestClient(app).get('/metrics')
|
||||
assert response.status_code == 200
|
||||
assert 'sglang:diffusion_requests_total' in response.text
|
||||
args.enable_metrics = False
|
||||
assert TestClient(create_app(args)).get('/metrics').status_code == 404
|
||||
"""
|
||||
subprocess.run([sys.executable, "-c", scrape], env=env, check=True, timeout=90)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, *sys.argv[1:]]))
|
||||
@@ -5,6 +5,7 @@ from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from prometheus_client import CollectorRegistry
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||
from sglang.multimodal_gen.runtime.layers.lora.linear import (
|
||||
@@ -13,6 +14,8 @@ from sglang.multimodal_gen.runtime.layers.lora.linear import (
|
||||
wrap_with_lora_layer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
|
||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
||||
from sglang.multimodal_gen.runtime.observability.metrics import DiffusionMetrics
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_lora
|
||||
|
||||
@@ -54,6 +57,37 @@ def _make_pipeline(layer: BaseLayerWithLoRA) -> _TestLoRAPipeline:
|
||||
return pipeline
|
||||
|
||||
|
||||
def test_worker_metrics_count_individual_adapters_in_multi_lora():
|
||||
pipeline = _make_pipeline(_make_layer())
|
||||
pipeline._temporarily_disable_offload = lambda *args, **kwargs: nullcontext([])
|
||||
pipeline.loaded_adapter_paths["second"] = "/second"
|
||||
pipeline.loaded_adapter_alphas["second"] = None
|
||||
pipeline.lora_adapters["second"] = pipeline.lora_adapters["adapter"]
|
||||
registry = CollectorRegistry()
|
||||
worker = GPUWorker.__new__(GPUWorker)
|
||||
worker.pipeline = pipeline
|
||||
worker.metrics = DiffusionMetrics(role="monolithic", replica="0", registry=registry)
|
||||
with patch(_RANK_PATCH, return_value=0):
|
||||
worker.set_lora(
|
||||
["adapter", "second"],
|
||||
[None, None],
|
||||
target="transformer",
|
||||
strength=[0.5, 0.5],
|
||||
merge_mode="merge",
|
||||
)
|
||||
assert (
|
||||
registry.get_sample_value(
|
||||
"sglang:diffusion_lora_active_adapters",
|
||||
{"role": "monolithic", "replica": "0"},
|
||||
)
|
||||
== 2
|
||||
)
|
||||
assert pipeline.get_lora_status()["active"]["transformer"][0]["nicknames"] == [
|
||||
"adapter",
|
||||
"second",
|
||||
]
|
||||
|
||||
|
||||
def test_merge_cache_only_accepts_cpu_backed_weights():
|
||||
pipeline = _make_pipeline(_make_layer())
|
||||
cpu_cache = pipeline._merge_cache_for(
|
||||
|
||||
Reference in New Issue
Block a user