feat(metrics): add Prometheus metrics for the EPD encoder server (#27564)
This commit is contained in:
@@ -56,6 +56,7 @@ from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
||||
from sglang.srt.model_loader import get_model
|
||||
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
|
||||
from sglang.srt.observability.metrics_collector import EncoderMetricsCollector
|
||||
from sglang.srt.observability.req_time_stats import EncoderReqTimeStats
|
||||
from sglang.srt.observability.trace import (
|
||||
process_tracing_init,
|
||||
@@ -67,11 +68,13 @@ from sglang.srt.server_args import (
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
add_prometheus_middleware,
|
||||
configure_logger,
|
||||
load_audio,
|
||||
load_image,
|
||||
load_video,
|
||||
random_uuid,
|
||||
set_prometheus_multiproc_dir,
|
||||
)
|
||||
from sglang.srt.utils.common import configure_logger, maybe_reindex_device_id
|
||||
from sglang.srt.utils.network import (
|
||||
@@ -86,6 +89,11 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
HEALTH_CHECK_TIMEOUT = 30
|
||||
|
||||
|
||||
def is_health_check_request(rid: Optional[str]) -> bool:
|
||||
return isinstance(rid, str) and rid.startswith(HEALTH_CHECK_RID_PREFIX)
|
||||
|
||||
|
||||
# Minimal 32x32 black PNG for health check dummy encode
|
||||
MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
|
||||
|
||||
@@ -242,6 +250,9 @@ class MMEncoder:
|
||||
self.server_args = server_args
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
self.rank = rank
|
||||
# DP rank for metric labels; overridden by run_dp_worker in DP mode.
|
||||
# 0 in the single-instance (non-DP) path.
|
||||
self.dp_rank = 0
|
||||
self.profiler = EncoderProfiler(rank)
|
||||
self._load_mm_processor(server_args)
|
||||
|
||||
@@ -844,12 +855,17 @@ class MMEncoder:
|
||||
else:
|
||||
mm_item.set(k, val)
|
||||
|
||||
forward_start = time.perf_counter()
|
||||
with torch.inference_mode():
|
||||
new_embeddings = get_feature_fn([mm_item])
|
||||
if not keep_on_gpu:
|
||||
new_embeddings = new_embeddings.cpu()
|
||||
if new_embeddings.ndim != 2:
|
||||
new_embeddings = new_embeddings.reshape(-1, new_embeddings.shape[-1])
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.observe_model_forward(
|
||||
time.perf_counter() - forward_start, modality=modality.name.lower()
|
||||
)
|
||||
|
||||
sub_grids = [grid_thw[i] for i in indices]
|
||||
return self.slice_embedding(new_embeddings, sub_grids, modality)
|
||||
@@ -1427,9 +1443,10 @@ class MMEncoder:
|
||||
|
||||
return normalized
|
||||
|
||||
async def _process_mm_items(self, mm_items, modality):
|
||||
async def _process_mm_items(self, mm_items, modality, log_metrics: bool = True):
|
||||
model_preprocessor = getattr(self.model, "preprocess_mm_for_encoder", None)
|
||||
|
||||
preprocess_start = time.perf_counter()
|
||||
if modality == Modality.IMAGE:
|
||||
processor_input = await self._process_image_items(
|
||||
mm_items, model_preprocessor
|
||||
@@ -1444,6 +1461,10 @@ class MMEncoder:
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported modality: {modality}")
|
||||
if encoder_metrics_collector is not None and log_metrics:
|
||||
encoder_metrics_collector.observe_preprocess(
|
||||
time.perf_counter() - preprocess_start, modality=modality.name.lower()
|
||||
)
|
||||
|
||||
target = self.model.thinker if hasattr(self.model, "thinker") else self.model
|
||||
get_feature_method = getattr(target, f"get_{modality.name.lower()}_feature")
|
||||
@@ -1555,9 +1576,16 @@ class MMEncoder:
|
||||
processor_input["audio_feature_lens"] = output_lengths
|
||||
return processor_input
|
||||
|
||||
async def _encode(self, mm_items, modality: Modality) -> torch.Tensor:
|
||||
async def _encode(
|
||||
self, mm_items, modality: Modality, log_metrics: bool = True
|
||||
) -> torch.Tensor:
|
||||
modality_str = modality.name.lower()
|
||||
try:
|
||||
mm_inputs, get_feature_fn = await self._process_mm_items(mm_items, modality)
|
||||
# preprocess latency is observed inside _process_mm_items so all
|
||||
# callers (encode / batch_encode / global-cache) are covered.
|
||||
mm_inputs, get_feature_fn = await self._process_mm_items(
|
||||
mm_items, modality, log_metrics=log_metrics
|
||||
)
|
||||
except NotImplementedError as e:
|
||||
raise InternalError(f"Not implemented error: {str(e)}")
|
||||
except Exception as e:
|
||||
@@ -1578,24 +1606,58 @@ class MMEncoder:
|
||||
continue
|
||||
mm_item.set(k, _convert(v))
|
||||
|
||||
if self.server_args.enable_prefix_mm_cache:
|
||||
cache_hit = False
|
||||
use_mm_cache = self.server_args.enable_prefix_mm_cache and log_metrics
|
||||
if use_mm_cache:
|
||||
mm_item.set_pad_value()
|
||||
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
|
||||
async with self.mm_cache_lock:
|
||||
mm_cache = self.mm_cache.get([mm_item.hash])
|
||||
if mm_cache is not None:
|
||||
mm_embedding = mm_cache.embedding
|
||||
cache_hit = True
|
||||
|
||||
if mm_embedding is None:
|
||||
forward_start = time.perf_counter()
|
||||
with torch.inference_mode():
|
||||
mm_embedding: torch.Tensor = get_feature_fn([mm_item])
|
||||
mm_embedding = mm_embedding.cpu()
|
||||
if len(mm_embedding.shape) != 2:
|
||||
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
|
||||
if encoder_metrics_collector is not None and log_metrics:
|
||||
encoder_metrics_collector.observe_model_forward(
|
||||
time.perf_counter() - forward_start, modality=modality_str
|
||||
)
|
||||
|
||||
if self.server_args.enable_prefix_mm_cache:
|
||||
# Per-request cache hit metrics: tokens = embedding rows, files = 1 item.
|
||||
if use_mm_cache and encoder_metrics_collector is not None:
|
||||
total_tokens = int(mm_embedding.shape[0])
|
||||
hit_tokens = total_tokens if cache_hit else 0
|
||||
encoder_metrics_collector.record_cache_tokens(
|
||||
hit_tokens, total_tokens, modality=modality_str
|
||||
)
|
||||
encoder_metrics_collector.record_cache_files(
|
||||
1 if cache_hit else 0, 1, modality=modality_str
|
||||
)
|
||||
|
||||
if use_mm_cache:
|
||||
async with self.mm_cache_lock:
|
||||
self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding))
|
||||
entries_before = len(self.mm_cache)
|
||||
already_present = self.mm_cache.has(mm_hash)
|
||||
inserted = self.mm_cache.set(
|
||||
mm_hash, EmbeddingResult(embedding=mm_embedding)
|
||||
)
|
||||
entries_after = len(self.mm_cache)
|
||||
if encoder_metrics_collector is not None:
|
||||
added = 0 if already_present else (1 if inserted else 0)
|
||||
evictions = max(0, added - (entries_after - entries_before))
|
||||
if evictions > 0:
|
||||
encoder_metrics_collector.inc_cache_evictions(
|
||||
modality=modality_str, count=evictions
|
||||
)
|
||||
encoder_metrics_collector.set_cache_state(
|
||||
self.mm_cache.current_size, entries_after
|
||||
)
|
||||
if self.profiler is not None:
|
||||
self.profiler.step()
|
||||
|
||||
@@ -1607,7 +1669,12 @@ class MMEncoder:
|
||||
)
|
||||
encode_video_audio_fn = getattr(target, "encode_video_audio", None)
|
||||
if encode_video_audio_fn is not None:
|
||||
audio_forward_start = time.perf_counter()
|
||||
audio_embedding = encode_video_audio_fn(mm_inputs)
|
||||
if encoder_metrics_collector is not None and log_metrics:
|
||||
encoder_metrics_collector.observe_model_forward(
|
||||
time.perf_counter() - audio_forward_start, modality="audio"
|
||||
)
|
||||
if audio_embedding is not None:
|
||||
aux_data["video_audio_embedding"] = audio_embedding
|
||||
else:
|
||||
@@ -1683,6 +1750,10 @@ class MMEncoder:
|
||||
embedding.nbytes,
|
||||
)
|
||||
xfer_ms = (time.monotonic() - _t_xfer_start) * 1000.0
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.observe_transfer(
|
||||
xfer_ms / 1000.0, backend="mooncake"
|
||||
)
|
||||
if not mr_already_registered:
|
||||
self.engine.deregister(embedding.data_ptr())
|
||||
# Only emit at INFO when transfer is slow or fell back
|
||||
@@ -1731,13 +1802,25 @@ class MMEncoder:
|
||||
finally:
|
||||
sock.close(linger=5000)
|
||||
|
||||
_zmq_xfer_start = time.perf_counter()
|
||||
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
|
||||
if (
|
||||
encoder_metrics_collector is not None
|
||||
and self.server_args.encoder_transfer_backend != "mooncake"
|
||||
):
|
||||
encoder_metrics_collector.observe_transfer(
|
||||
time.perf_counter() - _zmq_xfer_start,
|
||||
backend=self.server_args.encoder_transfer_backend,
|
||||
)
|
||||
|
||||
async def encode(
|
||||
self, mm_items, modality: Modality, req_id, num_parts, part_idx, hashes=None
|
||||
):
|
||||
try:
|
||||
grid_dim, mm_embedding, aux_data = await self._encode(mm_items, modality)
|
||||
log_metrics = not is_health_check_request(req_id)
|
||||
grid_dim, mm_embedding, aux_data = await self._encode(
|
||||
mm_items, modality, log_metrics=log_metrics
|
||||
)
|
||||
|
||||
if self.rank == 0:
|
||||
mm_data = EmbeddingData(
|
||||
@@ -1995,6 +2078,16 @@ class MMEncoder:
|
||||
items_per_req.append(sum(self._grid_count_per_leaf(leaves, modality)))
|
||||
total = sum(items_per_req)
|
||||
|
||||
if encoder_metrics_collector is not None:
|
||||
modality_str = modality.name.lower()
|
||||
for n in items_per_req:
|
||||
encoder_metrics_collector.observe_mm_items_per_request(
|
||||
n, modality=modality_str
|
||||
)
|
||||
encoder_metrics_collector.observe_mm_items_per_batch(
|
||||
total, modality=modality_str
|
||||
)
|
||||
|
||||
try:
|
||||
mm_inputs, get_feat = await self._process_mm_items(flat_items, modality)
|
||||
except NotImplementedError as e:
|
||||
@@ -2396,6 +2489,12 @@ class EncoderScheduler:
|
||||
|
||||
requests = [p.request for p in group]
|
||||
start = time.time()
|
||||
modality_str = modality.name.lower()
|
||||
if encoder_metrics_collector is not None:
|
||||
for p in group:
|
||||
encoder_metrics_collector.observe_queue_wait(
|
||||
max(0.0, start - p.submit_time), modality=modality_str
|
||||
)
|
||||
for sock in self.send_sockets:
|
||||
sock_send(
|
||||
sock,
|
||||
@@ -2448,9 +2547,26 @@ class EncoderScheduler:
|
||||
group: List[PendingRequest],
|
||||
modality: Modality,
|
||||
) -> None:
|
||||
modality_str = modality.name.lower()
|
||||
for p in group:
|
||||
req = p.request
|
||||
try:
|
||||
start = time.time()
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.observe_queue_wait(
|
||||
max(0.0, start - p.submit_time), modality=modality_str
|
||||
)
|
||||
# Count like batch_encode: flatten nested items and expand
|
||||
# per-leaf grids so {"type": "image", "image": [p1, p2, ...]}
|
||||
# counts as N, not 1.
|
||||
leaves = MMEncoder._flatten_nested_items(req.get("mm_items", []))
|
||||
mm_count = sum(self.encoder._grid_count_per_leaf(leaves, modality))
|
||||
encoder_metrics_collector.observe_mm_items_per_request(
|
||||
mm_count, modality=modality_str
|
||||
)
|
||||
encoder_metrics_collector.observe_mm_items_per_batch(
|
||||
mm_count, modality=modality_str
|
||||
)
|
||||
for sock in self.send_sockets:
|
||||
sock_send(sock, wrap_as_pickle(req))
|
||||
result = await self.encoder.encode_request(req, modality)
|
||||
@@ -2468,6 +2584,10 @@ encoder: Optional[MMEncoder] = None
|
||||
send_sockets: List[zmq.Socket] = []
|
||||
encoder_scheduler: Optional[EncoderScheduler] = None
|
||||
|
||||
# Per-process encoder metrics collector. Set in launch_server (non-DP) and in
|
||||
# run_dp_worker (DP mode, with the worker's dp_rank). None when metrics disabled.
|
||||
encoder_metrics_collector: Optional[EncoderMetricsCollector] = None
|
||||
|
||||
# DP mode (--dp-size > 1): each rank runs as a subprocess with its own
|
||||
# MMEncoder on its own GPU; the main process only routes via ZMQ so the
|
||||
# asyncio event loop is never blocked by GPU work.
|
||||
@@ -2519,6 +2639,8 @@ async def _dp_worker_encode_and_send(
|
||||
time_stats.decode_json(time_stats_json)
|
||||
request["enter_time"] = time.time()
|
||||
modality = Modality.from_str(request["modality"])
|
||||
time_stats.modality = modality.name.lower()
|
||||
time_stats.set_metrics_collector(encoder_metrics_collector)
|
||||
backend = enc.server_args.encoder_transfer_backend
|
||||
|
||||
# URL state lives in main process module globals; workers don't see it.
|
||||
@@ -2638,6 +2760,8 @@ class DPDispatcher:
|
||||
dispatch_sockets: List,
|
||||
result_socket,
|
||||
worker_processes: List[mp.Process],
|
||||
enable_metrics: bool = False,
|
||||
labels: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
self.dp_size = dp_size
|
||||
self.dispatch_sockets = dispatch_sockets
|
||||
@@ -2656,10 +2780,30 @@ class DPDispatcher:
|
||||
# Set when _result_listener gives up; makes alive_ranks report empty.
|
||||
self._listener_failed = False
|
||||
|
||||
# Prometheus gauge: pending requests per DP rank. Lives in the main
|
||||
# process (the dispatcher), unlike the per-worker EncoderMetricsCollector.
|
||||
self.labels = dict(labels or {})
|
||||
self.pending_gauge = None
|
||||
if enable_metrics:
|
||||
from prometheus_client import Gauge
|
||||
|
||||
self.pending_gauge = Gauge(
|
||||
name="sglang:encoder_dp_pending_requests",
|
||||
documentation="Number of pending requests per encoder DP rank.",
|
||||
labelnames=list(self.labels.keys()) + ["dp_rank"],
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
|
||||
@property
|
||||
def pending_counts(self) -> List[int]:
|
||||
return [len(d) for d in self.pending_futures]
|
||||
|
||||
def _update_pending_gauge(self) -> None:
|
||||
"""Push current pending counts to the Prometheus gauge (absolute set)."""
|
||||
if self.pending_gauge is not None:
|
||||
for i, c in enumerate(self.pending_counts):
|
||||
self.pending_gauge.labels(**self.labels, dp_rank=str(i)).set(c)
|
||||
|
||||
@property
|
||||
def alive_ranks(self) -> List[int]:
|
||||
# Empty if the result listener died; else ranks not marked dead.
|
||||
@@ -2682,6 +2826,7 @@ class DPDispatcher:
|
||||
# dispatch / broadcast failure: no follow-up /send expected.
|
||||
self.pending_futures[rank].pop(req_id, None)
|
||||
self.req_id_to_rank.pop(req_id, None)
|
||||
self._update_pending_gauge()
|
||||
|
||||
def _fail_pending_for_rank(self, rank: int, reason: str, error_type: str) -> None:
|
||||
# Resolve a rank's outstanding futures with 503 so awaiters don't hang.
|
||||
@@ -2699,6 +2844,7 @@ class DPDispatcher:
|
||||
}
|
||||
)
|
||||
pending.pop(key, None)
|
||||
self._update_pending_gauge()
|
||||
|
||||
def _fail_all_pending(self, reason: str, error_type: str) -> None:
|
||||
for rank in range(self.dp_size):
|
||||
@@ -2734,6 +2880,7 @@ class DPDispatcher:
|
||||
self.req_id_to_rank[req_id] = rank
|
||||
future = asyncio.get_running_loop().create_future()
|
||||
self.pending_futures[rank][req_id] = future
|
||||
self._update_pending_gauge()
|
||||
logger.info(
|
||||
f"MM-Encoder DP dispatch: req_id={req_id}, "
|
||||
f"modality={request.get('modality', 'image')}, "
|
||||
@@ -2944,6 +3091,7 @@ class DPDispatcher:
|
||||
)
|
||||
continue
|
||||
future = self.pending_futures[rank].pop(key)
|
||||
self._update_pending_gauge()
|
||||
# Only mooncake encode (content=request dict) needs the mapping
|
||||
# kept for the follow-up /send.
|
||||
keep_mapping = dp_type == "encode" and msg.get("content") is not None
|
||||
@@ -3011,6 +3159,15 @@ async def _dp_worker_handle_request(
|
||||
dp_type: str,
|
||||
) -> None:
|
||||
t0 = time.time()
|
||||
modality_str = str(request.get("modality", "image")).lower()
|
||||
is_encode = dp_type not in (
|
||||
"start_profile",
|
||||
"stop_profile",
|
||||
"health_encode",
|
||||
"send",
|
||||
)
|
||||
if is_encode and encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_received(modality=modality_str)
|
||||
try:
|
||||
if dp_type in ("start_profile", "stop_profile"):
|
||||
content = await _dp_worker_handle_profile(enc, dp_rank, dp_type, request)
|
||||
@@ -3037,6 +3194,10 @@ async def _dp_worker_handle_request(
|
||||
f"modality={request.get('modality', 'image')}, "
|
||||
f"cost={(time.time() - t0) * 1000:.1f}ms"
|
||||
)
|
||||
if is_encode and encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="success"
|
||||
)
|
||||
envelope = {
|
||||
"req_id": request.get("req_id", ""),
|
||||
"_dp_type": dp_type,
|
||||
@@ -3048,6 +3209,10 @@ async def _dp_worker_handle_request(
|
||||
f"req_id={request.get('req_id', '?')}: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
if is_encode and encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="error"
|
||||
)
|
||||
err_code = int(getattr(e, "code", None) or HTTPStatus.INTERNAL_SERVER_ERROR)
|
||||
envelope = {
|
||||
"req_id": request.get("req_id", ""),
|
||||
@@ -3089,6 +3254,19 @@ async def run_dp_worker(
|
||||
args.base_gpu_id = gpu_id
|
||||
args.tp_size = 1
|
||||
enc = MMEncoder(args, dist_init_method=f"tcp://127.0.0.1:{get_free_port()}", rank=0)
|
||||
|
||||
global encoder_metrics_collector
|
||||
if server_args.enable_metrics:
|
||||
set_prometheus_multiproc_dir()
|
||||
labels = {
|
||||
"model_name": server_args.served_model_name,
|
||||
"dp_rank": str(dp_rank),
|
||||
}
|
||||
if server_args.extra_metric_labels:
|
||||
labels.update(server_args.extra_metric_labels)
|
||||
encoder_metrics_collector = EncoderMetricsCollector(labels)
|
||||
enc.dp_rank = dp_rank
|
||||
|
||||
sched = EncoderScheduler(
|
||||
encoder=enc, send_sockets=[], max_batch_size=ENCODER_MAX_BATCH_SIZE
|
||||
)
|
||||
@@ -3354,7 +3532,20 @@ def launch_server(server_args: ServerArgs):
|
||||
_launch_server_dp(server_args)
|
||||
return
|
||||
|
||||
global encoder
|
||||
global encoder, encoder_metrics_collector
|
||||
|
||||
# Set up prometheus metrics.
|
||||
if server_args.enable_metrics:
|
||||
set_prometheus_multiproc_dir()
|
||||
labels = {
|
||||
"model_name": server_args.served_model_name,
|
||||
"dp_rank": "0",
|
||||
}
|
||||
if server_args.extra_metric_labels:
|
||||
labels.update(server_args.extra_metric_labels)
|
||||
encoder_metrics_collector = EncoderMetricsCollector(labels)
|
||||
add_prometheus_middleware(app)
|
||||
|
||||
ctx = mp.get_context("spawn")
|
||||
zmq_ctx = zmq.Context(10)
|
||||
ipc_path_prefix = random_uuid()
|
||||
@@ -3406,6 +3597,12 @@ def _launch_server_dp(server_args: ServerArgs):
|
||||
dp_size = server_args.dp_size
|
||||
logger.info(f"Launching encoder in DP mode: dp_size={dp_size}")
|
||||
|
||||
# DP mode: workers (subprocesses) write metrics to the shared multiproc dir;
|
||||
# the main process exposes the aggregated /metrics endpoint.
|
||||
if server_args.enable_metrics:
|
||||
set_prometheus_multiproc_dir()
|
||||
add_prometheus_middleware(app)
|
||||
|
||||
ctx = mp.get_context("spawn")
|
||||
ipc_prefix = random_uuid()
|
||||
async_zmq_ctx = zmq.asyncio.Context(dp_size + 1)
|
||||
@@ -3457,11 +3654,16 @@ def _launch_server_dp(server_args: ServerArgs):
|
||||
proc.start()
|
||||
worker_processes.append(proc)
|
||||
|
||||
labels = {"model_name": server_args.served_model_name}
|
||||
if server_args.extra_metric_labels:
|
||||
labels.update(server_args.extra_metric_labels)
|
||||
dp_dispatcher = DPDispatcher(
|
||||
dp_size,
|
||||
dispatch_sockets,
|
||||
result_socket,
|
||||
worker_processes,
|
||||
enable_metrics=server_args.enable_metrics,
|
||||
labels=labels,
|
||||
)
|
||||
|
||||
# Register this encoder's URL with prefill server(s) if configured.
|
||||
@@ -3551,6 +3753,8 @@ async def handle_encode_request(request: dict):
|
||||
f"modality={request.get('modality', 'image')}"
|
||||
)
|
||||
return ORJSONResponse(content=result.get("content"))
|
||||
|
||||
modality_str = str(request.get("modality", "image")).lower()
|
||||
try:
|
||||
# when multiple decoder TP ranks POST /encode
|
||||
# with the same req_id, only the first triggers the VIT forward;
|
||||
@@ -3603,7 +3807,12 @@ async def handle_encode_request(request: dict):
|
||||
if time_stats_json:
|
||||
time_stats.decode_json(time_stats_json)
|
||||
|
||||
modality_str = modality.name.lower()
|
||||
time_stats.modality = modality_str
|
||||
time_stats.set_metrics_collector(encoder_metrics_collector)
|
||||
time_stats.set_mm_encode_start_time()
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_received(modality=modality_str)
|
||||
if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES:
|
||||
try:
|
||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||
@@ -3651,6 +3860,10 @@ async def handle_encode_request(request: dict):
|
||||
if evt:
|
||||
evt.set()
|
||||
await encoder._cleanup_inflight_encode_state(req_id)
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="error"
|
||||
)
|
||||
return ORJSONResponse(
|
||||
status_code=error_code,
|
||||
content={"status": "error", "message": error_msg, "req_id": req_id},
|
||||
@@ -3674,6 +3887,10 @@ async def handle_encode_request(request: dict):
|
||||
"embedding_dim": embedding_dim,
|
||||
}
|
||||
)
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="success"
|
||||
)
|
||||
return ORJSONResponse(content=request)
|
||||
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
|
||||
logger.info(f"{request['embedding_port'] = }")
|
||||
@@ -3694,6 +3911,10 @@ async def handle_encode_request(request: dict):
|
||||
)
|
||||
await asyncio.gather(*tasks)
|
||||
encoder.embedding_to_send.pop(request["req_id"], None)
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="success"
|
||||
)
|
||||
return ORJSONResponse(content=None)
|
||||
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
|
||||
await encoder.send(
|
||||
@@ -3707,6 +3928,10 @@ async def handle_encode_request(request: dict):
|
||||
f"[{req_id}] /encode completed in {elapsed:.3f}s, "
|
||||
f"modality={request['modality']}, tokens={embedding_len}"
|
||||
)
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="success"
|
||||
)
|
||||
return ORJSONResponse(content=None)
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
@@ -3719,6 +3944,10 @@ async def handle_encode_request(request: dict):
|
||||
if evt:
|
||||
evt.set()
|
||||
await encoder._cleanup_inflight_encode_state(req_id)
|
||||
if encoder_metrics_collector is not None:
|
||||
encoder_metrics_collector.inc_requests_total(
|
||||
modality=modality_str, status="error"
|
||||
)
|
||||
return ORJSONResponse(
|
||||
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
|
||||
content={
|
||||
|
||||
@@ -1942,6 +1942,238 @@ class RadixCacheMetricsCollector(_StatLoggerDIMixin):
|
||||
self.load_back_duration_seconds.labels(**self.labels).observe(duration_seconds)
|
||||
|
||||
|
||||
class EncoderMetricsCollector(_StatLoggerDIMixin):
|
||||
"""Metrics collector for the EPD encoder server (--encoder-only)."""
|
||||
|
||||
def __init__(self, labels: Dict[str, str]) -> None:
|
||||
# We need to import prometheus_client after setting the env variable `PROMETHEUS_MULTIPROC_DIR`
|
||||
from prometheus_client import Counter as _PromCounter
|
||||
from prometheus_client import Gauge as _PromGauge
|
||||
from prometheus_client import Histogram as _PromHistogram
|
||||
|
||||
Counter = self._counter_cls or _PromCounter
|
||||
Gauge = self._gauge_cls or _PromGauge
|
||||
Histogram = self._histogram_cls or _PromHistogram
|
||||
|
||||
self.labels = labels
|
||||
|
||||
self.cache_evictions_total = Counter(
|
||||
name="sglang:encoder_cache_evictions_total",
|
||||
documentation="Total cache evictions.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
)
|
||||
self.cache_size_mb = Gauge(
|
||||
name="sglang:encoder_cache_size_mb",
|
||||
documentation="Current cache size in MB.",
|
||||
labelnames=labels.keys(),
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
self.cache_entries = Gauge(
|
||||
name="sglang:encoder_cache_entries",
|
||||
documentation="Current number of cache entries.",
|
||||
labelnames=labels.keys(),
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
self.cache_hit_tokens_total = Counter(
|
||||
name="sglang:encoder_cache_hit_tokens_total",
|
||||
documentation="Total tokens served from cache (cache hits).",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
)
|
||||
self.cache_total_tokens_total = Counter(
|
||||
name="sglang:encoder_cache_total_tokens_total",
|
||||
documentation="Total tokens processed (hit + miss).",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
)
|
||||
self.cache_hit_files_total = Counter(
|
||||
name="sglang:encoder_cache_hit_files_total",
|
||||
documentation="Total files served from cache.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
)
|
||||
self.cache_total_files_total = Counter(
|
||||
name="sglang:encoder_cache_total_files_total",
|
||||
documentation="Total files processed (hit + miss).",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
)
|
||||
|
||||
# Total encoder requests by modality and status
|
||||
self.requests_total = Counter(
|
||||
name="sglang:encoder_requests_total",
|
||||
documentation="Total encoder requests by modality and status.",
|
||||
labelnames=list(labels.keys()) + ["modality", "status"],
|
||||
)
|
||||
|
||||
# Total requests received per DP rank (incremented at receive time, before processing).
|
||||
# Use rate(sglang:encoder_requests_received_total[1m]) for per-encoder QPS.
|
||||
self.requests_received_total = Counter(
|
||||
name="sglang:encoder_requests_received_total",
|
||||
documentation="Total requests received by encoder (at receive time), per DP rank.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
)
|
||||
|
||||
# Multimodal items per batch histogram
|
||||
self.mm_items_per_batch = Histogram(
|
||||
name="sglang:encoder_mm_items_per_batch",
|
||||
documentation="Histogram of multimodal items processed per encoder batch.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
buckets=[
|
||||
1,
|
||||
2,
|
||||
3,
|
||||
4,
|
||||
5,
|
||||
6,
|
||||
7,
|
||||
8,
|
||||
9,
|
||||
10,
|
||||
11,
|
||||
12,
|
||||
13,
|
||||
14,
|
||||
15,
|
||||
16,
|
||||
32,
|
||||
64,
|
||||
128,
|
||||
],
|
||||
)
|
||||
|
||||
# Multimodal items per request histogram
|
||||
self.mm_items_per_request = Histogram(
|
||||
name="sglang:encoder_mm_items_per_request",
|
||||
documentation="Histogram of multimodal items per individual encoder request.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
buckets=[1, 2, 3, 4, 5, 6, 7, 8, 10, 12, 16, 24, 32, 64],
|
||||
)
|
||||
|
||||
# Per-request E2E encoder latency
|
||||
self.encoder_request_e2e_latency_seconds = Histogram(
|
||||
name="sglang:encoder_request_e2e_latency_seconds",
|
||||
documentation="Histogram of per-request end-to-end encoder latency in seconds (queue wait + encode).",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
buckets=[0.01, 0.02, 0.05, 0.1, 0.2, 0.5, 1, 2, 5, 10, 20, 30, 60],
|
||||
)
|
||||
|
||||
# --- Latency breakdown histograms ---
|
||||
|
||||
# Queue wait: time spent in scheduler queue before batch processing starts
|
||||
self.queue_wait_seconds = Histogram(
|
||||
name="sglang:encoder_queue_wait_seconds",
|
||||
documentation="Time request spent waiting in scheduler queue.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
buckets=[0.001, 0.005, 0.01, 0.05, 0.1, 0.5, 1, 2, 5, 10],
|
||||
)
|
||||
|
||||
# Preprocess: CPU data loading + processor (image decode, video frame sampling, etc.)
|
||||
self.preprocess_seconds = Histogram(
|
||||
name="sglang:encoder_preprocess_seconds",
|
||||
documentation="Data loading and preprocessing latency.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
buckets=[0.01, 0.05, 0.1, 0.2, 0.5, 1, 2, 5, 10, 30],
|
||||
)
|
||||
|
||||
# Model forward: model forward pass latency
|
||||
self.model_forward_seconds = Histogram(
|
||||
name="sglang:encoder_model_forward_seconds",
|
||||
documentation="GPU model forward pass latency.",
|
||||
labelnames=list(labels.keys()) + ["modality"],
|
||||
buckets=[0.01, 0.02, 0.05, 0.1, 0.2, 0.5, 1, 2, 5],
|
||||
)
|
||||
|
||||
# Embedding transfer: embedding transfer to prefill node (zmq or mooncake)
|
||||
self.transfer_seconds = Histogram(
|
||||
name="sglang:encoder_transfer_seconds",
|
||||
documentation="Embedding transfer latency to prefill node.",
|
||||
labelnames=list(labels.keys()) + ["backend"],
|
||||
buckets=[0.001, 0.005, 0.01, 0.05, 0.1, 0.2, 0.5, 1, 2],
|
||||
)
|
||||
|
||||
def _inc_cache_counter(self, counter, modality: str, count: int = 1) -> None:
|
||||
counter.labels(**self.labels, modality=modality).inc(count)
|
||||
|
||||
def inc_cache_evictions(self, modality: str = "image", count: int = 1) -> None:
|
||||
self._inc_cache_counter(self.cache_evictions_total, modality, count)
|
||||
|
||||
def record_cache_tokens(
|
||||
self, hit_tokens: int, total_tokens: int, modality: str = "image"
|
||||
) -> None:
|
||||
self._inc_cache_counter(self.cache_total_tokens_total, modality, total_tokens)
|
||||
if hit_tokens > 0:
|
||||
self._inc_cache_counter(self.cache_hit_tokens_total, modality, hit_tokens)
|
||||
|
||||
def record_cache_files(
|
||||
self, hit_files: int, total_files: int, modality: str = "image"
|
||||
) -> None:
|
||||
self._inc_cache_counter(self.cache_total_files_total, modality, total_files)
|
||||
if hit_files > 0:
|
||||
self._inc_cache_counter(self.cache_hit_files_total, modality, hit_files)
|
||||
|
||||
def set_cache_state(self, current_size: int, num_entries: int) -> None:
|
||||
self.cache_size_mb.labels(**self.labels).set(current_size / (1024 * 1024))
|
||||
self.cache_entries.labels(**self.labels).set(num_entries)
|
||||
|
||||
def observe_queue_wait(
|
||||
self, latency_seconds: float, modality: str = "image"
|
||||
) -> None:
|
||||
"""Record time spent waiting in the scheduler queue."""
|
||||
self.queue_wait_seconds.labels(**self.labels, modality=modality).observe(
|
||||
latency_seconds
|
||||
)
|
||||
|
||||
def observe_preprocess(
|
||||
self, latency_seconds: float, modality: str = "image"
|
||||
) -> None:
|
||||
"""Record data loading and preprocessing latency."""
|
||||
self.preprocess_seconds.labels(**self.labels, modality=modality).observe(
|
||||
latency_seconds
|
||||
)
|
||||
|
||||
def observe_model_forward(
|
||||
self, latency_seconds: float, modality: str = "image"
|
||||
) -> None:
|
||||
"""Record model forward pass latency."""
|
||||
self.model_forward_seconds.labels(**self.labels, modality=modality).observe(
|
||||
latency_seconds
|
||||
)
|
||||
|
||||
def observe_transfer(self, latency_seconds: float, backend: str = "zmq") -> None:
|
||||
"""Record embedding transfer latency."""
|
||||
self.transfer_seconds.labels(**self.labels, backend=backend).observe(
|
||||
latency_seconds
|
||||
)
|
||||
|
||||
def observe_mm_items_per_batch(self, count: int, modality: str = "image") -> None:
|
||||
"""Record the number of multimodal items processed in a batch."""
|
||||
self.mm_items_per_batch.labels(**self.labels, modality=modality).observe(count)
|
||||
|
||||
def observe_mm_items_per_request(self, count: int, modality: str = "image") -> None:
|
||||
"""Record the number of multimodal items in a single request."""
|
||||
self.mm_items_per_request.labels(**self.labels, modality=modality).observe(
|
||||
count
|
||||
)
|
||||
|
||||
def inc_requests_total(self, modality: str, status: str) -> None:
|
||||
"""Increment encoder request counter. status: 'success' | 'error'."""
|
||||
self.requests_total.labels(
|
||||
**self.labels, modality=modality, status=status
|
||||
).inc()
|
||||
|
||||
def inc_requests_received(self, modality: str = "image") -> None:
|
||||
"""Increment the received-requests counter at request-arrival time.
|
||||
|
||||
dp_rank is supplied via self.labels (set per process at construction).
|
||||
"""
|
||||
self.requests_received_total.labels(**self.labels, modality=modality).inc()
|
||||
|
||||
def observe_request_e2e_latency(
|
||||
self, latency_seconds: float, modality: str = "image"
|
||||
) -> None:
|
||||
"""Record per-request end-to-end encoder latency in seconds."""
|
||||
self.encoder_request_e2e_latency_seconds.labels(
|
||||
**self.labels, modality=modality
|
||||
).observe(latency_seconds)
|
||||
|
||||
|
||||
def get_histogram_conf_from_env(env_var_name: str) -> Optional[List[float]]:
|
||||
"""
|
||||
Get the histogram configuration from the environment variable.
|
||||
|
||||
@@ -26,6 +26,7 @@ from typing_extensions import Self
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.observability.metrics_collector import (
|
||||
EncoderMetricsCollector,
|
||||
SchedulerMetricsCollector,
|
||||
TokenizerMetricsCollector,
|
||||
)
|
||||
@@ -224,7 +225,11 @@ class RequestStage:
|
||||
class ReqTimeStatsBase:
|
||||
enable_metrics: bool = False
|
||||
metrics_collector: Optional[
|
||||
Union[SchedulerMetricsCollector, TokenizerMetricsCollector]
|
||||
Union[
|
||||
SchedulerMetricsCollector,
|
||||
TokenizerMetricsCollector,
|
||||
EncoderMetricsCollector,
|
||||
]
|
||||
] = None
|
||||
trace_ctx: Union[TraceReqContext, TraceNullContext] = field(
|
||||
default_factory=TraceNullContext
|
||||
@@ -258,7 +263,12 @@ class ReqTimeStatsBase:
|
||||
return "unknown"
|
||||
|
||||
def set_metrics_collector(
|
||||
self, collector: Union[SchedulerMetricsCollector, TokenizerMetricsCollector]
|
||||
self,
|
||||
collector: Union[
|
||||
SchedulerMetricsCollector,
|
||||
TokenizerMetricsCollector,
|
||||
EncoderMetricsCollector,
|
||||
],
|
||||
):
|
||||
if collector:
|
||||
self.enable_metrics = True
|
||||
@@ -1172,6 +1182,7 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
|
||||
class EncoderReqTimeStats(ReqTimeStatsBase):
|
||||
mm_encode_start_time: float = 0.0
|
||||
mm_encode_end_time: float = 0.0
|
||||
modality: str = "image"
|
||||
|
||||
def set_mm_encode_start_time(self, ts=None):
|
||||
ts = ts or time.perf_counter()
|
||||
@@ -1194,6 +1205,10 @@ class EncoderReqTimeStats(ReqTimeStatsBase):
|
||||
convert_time_to_realtime_ns(ts),
|
||||
thread_finish_flag=True,
|
||||
)
|
||||
if self.enable_metrics:
|
||||
self.metrics_collector.observe_request_e2e_latency(
|
||||
ts - self.mm_encode_start_time, modality=self.modality
|
||||
)
|
||||
|
||||
|
||||
def set_schedule_time_batch(batch: ScheduleBatch):
|
||||
|
||||
Reference in New Issue
Block a user