diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 3c20cf791..4f2260707 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -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={ diff --git a/python/sglang/srt/observability/metrics_collector.py b/python/sglang/srt/observability/metrics_collector.py index d1a9cca04..cc4d2d6a1 100644 --- a/python/sglang/srt/observability/metrics_collector.py +++ b/python/sglang/srt/observability/metrics_collector.py @@ -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. diff --git a/python/sglang/srt/observability/req_time_stats.py b/python/sglang/srt/observability/req_time_stats.py index 0c654795d..27734d988 100644 --- a/python/sglang/srt/observability/req_time_stats.py +++ b/python/sglang/srt/observability/req_time_stats.py @@ -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): diff --git a/test/registered/observability/test_encoder_server_metrics.py b/test/registered/observability/test_encoder_server_metrics.py new file mode 100644 index 000000000..16fcc3273 --- /dev/null +++ b/test/registered/observability/test_encoder_server_metrics.py @@ -0,0 +1,101 @@ +"""Integration test: the EPD encoder server exports sglang:encoder_* metrics.""" + +import unittest +import uuid +from typing import Dict, List +from urllib.parse import urlparse + +import requests +import zmq +from prometheus_client.parser import text_string_to_metric_families +from prometheus_client.samples import Sample + +from sglang.srt.disaggregation.encode_server import MINIMUM_PNG_PICTURE_BASE64 +from sglang.srt.utils import kill_process_tree +from sglang.srt.utils.network import get_zmq_socket_on_host +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_SMALL_VLM_MODEL_NAME_FOR_TEST, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=180, stage="base-b", runner_config="1-gpu-small") + +_MODEL_NAME = DEFAULT_SMALL_VLM_MODEL_NAME_FOR_TEST + + +def _parse_prometheus_metrics(metrics_text: str) -> Dict[str, List[Sample]]: + result: Dict[str, List[Sample]] = {} + for family in text_string_to_metric_families(metrics_text): + for sample in family.samples: + result.setdefault(sample.name, []).append(sample) + return result + + +class TestEncoderServerMetrics(CustomTestCase): + def test_encoder_metrics_exported(self): + process = popen_launch_server( + _MODEL_NAME, + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--encoder-only", + "--trust-remote-code", + "--enable-metrics", + "--tp", + "1", + ], + ) + self.addCleanup(kill_process_tree, process.pid) + base_host = urlparse(DEFAULT_URL_FOR_TEST).hostname + context = zmq.Context() + recv_port, recv_socket = get_zmq_socket_on_host( + context, zmq.PULL, host=base_host + ) + try: + health = requests.get(f"{DEFAULT_URL_FOR_TEST}/health_generate") + self.assertEqual(health.status_code, 200) + + req_id = f"metrics-probe-{uuid.uuid4().hex}" + requests.post( + f"{DEFAULT_URL_FOR_TEST}/scheduler_receive_url", + json={ + "req_id": req_id, + "receive_url": f"{base_host}:{recv_port}", + "receive_count": 1, + }, + ) + response = requests.post( + f"{DEFAULT_URL_FOR_TEST}/encode", + json={ + "req_id": req_id, + "modality": "IMAGE", + "mm_items": [f"data:image/png;base64,{MINIMUM_PNG_PICTURE_BASE64}"], + "num_parts": 1, + "part_idx": 0, + "embedding_port": None, + }, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + ) + self.assertEqual(response.status_code, 200) + + metrics_response = requests.get(f"{DEFAULT_URL_FOR_TEST}/metrics") + self.assertEqual(metrics_response.status_code, 200) + metrics_text = metrics_response.text + + self.assertIn("sglang:encoder_requests_received_total", metrics_text) + self.assertIn(f'model_name="{_MODEL_NAME}"', metrics_text) + + metrics = _parse_prometheus_metrics(metrics_text) + received = metrics.get("sglang:encoder_requests_received_total", []) + self.assertGreater(sum(s.value for s in received), 0) + finally: + recv_socket.close() + context.term() + + +if __name__ == "__main__": + unittest.main()