diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index aebe03942..c5c3b11aa 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -41,13 +41,16 @@ def __init__( self.metric_client = MetricClient(get_shm_port_args().metric_port) self.id_gen = ReqIDGenerator() - self.pd_manager = PDManager(args) + self.pd_manager = PDManager(args, self.metric_client) self.req_id_to_out_inf: Dict[int, ReqStatus] = {} self.infos_queues = None # 这个需要延迟初始化,否则使用的loop不对 self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() self.running_request_count = 0 + self.pd_stage_waiting_request_counts = {"prefill": 0, "decode": 0} + for stage in self.pd_stage_waiting_request_counts: + self.metric_client.gauge_set("lightllm_pd_master_stage_waiting_requests", 0, labels={"stage": stage}) self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) @@ -354,6 +357,28 @@ async def raise_if_disconnected() -> None: except asyncio.TimeoutError: continue + def _change_pd_stage_waiting_requests(self, stage: str, delta: int) -> None: + self.pd_stage_waiting_request_counts[stage] += delta + self.metric_client.gauge_set( + "lightllm_pd_master_stage_waiting_requests", + self.pd_stage_waiting_request_counts[stage], + labels={"stage": stage}, + ) + + async def _wait_for_pd_stage( + self, + event: asyncio.Event, + request: Request, + timeout: float, + group_request_id: int, + stage: str, + ) -> None: + self._change_pd_stage_waiting_requests(stage, 1) + try: + await self._wait_for_event_or_disconnect(event, request, timeout, group_request_id, stage) + finally: + self._change_pd_stage_waiting_requests(stage, -1) + async def _log_req_header(self, request: Request, group_request_id: int): x_request_id = request.headers.get("X-Request-Id", "") x_session_id = request.headers.get("X-Session-Id", "") @@ -388,7 +413,7 @@ async def fetch_pd_stream( await p_node.websocket.send_bytes(pickle.dumps((ObjType.REQ, (prompt, sampling_params, multimodal_params)))) try: - await self._wait_for_event_or_disconnect( + await self._wait_for_pd_stage( prefill_prompt_ids_event, request, timeout=60, @@ -409,7 +434,7 @@ async def fetch_pd_stream( ) try: - await self._wait_for_event_or_disconnect( + await self._wait_for_pd_stage( up_status_event, request, timeout=180, @@ -431,13 +456,19 @@ async def fetch_pd_stream( first_token_gen = False needs_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 - prompt_cache_len_from_prefill = await self._wait_for_prefill_token_if_needed( - req_status=req_status, - request=request, - group_request_id=group_request_id, - needs_prefill_first_token=needs_prefill_first_token, - ready_kv_len=decode_node_info.ready_kv_len, - ) + if needs_prefill_first_token: + self._change_pd_stage_waiting_requests("prefill", 1) + try: + prompt_cache_len_from_prefill = await self._wait_for_prefill_token_if_needed( + req_status=req_status, + request=request, + group_request_id=group_request_id, + needs_prefill_first_token=needs_prefill_first_token, + ready_kv_len=decode_node_info.ready_kv_len, + ) + finally: + if needs_prefill_first_token: + self._change_pd_stage_waiting_requests("prefill", -1) while True: await req_status.wait_to_ready() @@ -749,8 +780,9 @@ async def put_tokens_to_front(self, token_list: List[Tuple[int, str, dict, Finis class PDManager: - def __init__(self, args: StartArgs): + def __init__(self, args: StartArgs, metric_client=None): self.args: StartArgs = args + self.metric_client = metric_client self.prefill_nodes: List[PD_Client_Obj] = [] self.decode_nodes: List[PD_Client_Obj] = [] self.url_to_pd_nodes: Dict[str, PD_Client_Obj] = {} @@ -874,6 +906,12 @@ def update_node_load_info(self, load_info: Optional[dict]): total_token_usage_rate = load_info["total_token_usage_rate"] pd_client = self.url_to_pd_nodes.get(client_ip_port) pd_client.run_status.total_token_usage_rate = total_token_usage_rate + if self.metric_client is not None: + self.metric_client.gauge_set( + "lightllm_pd_node_token_usage_ratio", + total_token_usage_rate, + labels={"role": pd_client.mode, "endpoint": client_ip_port}, + ) except BaseException as e: logger.warning(f"udpate node load info failed, load_info: {load_info} error: {str(e)}") return diff --git a/lightllm/server/metrics/manager.py b/lightllm/server/metrics/manager.py index 22f6426a7..f8f56dc8d 100644 --- a/lightllm/server/metrics/manager.py +++ b/lightllm/server/metrics/manager.py @@ -55,8 +55,9 @@ def exposed_counter_inc_by(self, name: str, amount: float) -> None: def exposed_histogram_observe(self, name: str, value: float, label: str = None) -> None: return self.monitor.histogram_observe(name, value, label) - def exposed_gauge_set(self, name: str, value: float) -> None: - return self.monitor.gauge_set(name, value) + def exposed_gauge_set(self, name: str, value: float, labels: dict = None) -> None: + local_labels = None if labels is None else {key: labels[key] for key in labels} + return self.monitor.gauge_set(name, value, local_labels) def exposed_generate_latest(self) -> bytes: data = generate_latest(self.monitor.registry) diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 0d42462c3..b58b587b2 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,6 +32,8 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", + "lightllm_pd_node_token_usage_ratio": "Token capacity usage ratio reported by a PD node", + "lightllm_pd_master_stage_waiting_requests": "Requests waiting for a PD stage to become ready", } @@ -111,6 +113,8 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") + self.create_gauge("lightllm_pd_node_token_usage_ratio", labelnames=["role", "endpoint"]) + self.create_gauge("lightllm_pd_master_stage_waiting_requests", labelnames=["stage"]) def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) @@ -122,8 +126,9 @@ def create_counter(self, name, labelnames=None): counter = Counter(name, MONITOR_INFO[name], labelnames=all_labels, registry=self.registry) self.monitor_registry[name] = counter - def create_gauge(self, name): - gauge = Gauge(name, MONITOR_INFO[name], labelnames=["model_name"], registry=self.registry) + def create_gauge(self, name, labelnames=None): + all_labels = ["model_name"] + (labelnames or []) + gauge = Gauge(name, MONITOR_INFO[name], labelnames=all_labels, registry=self.registry) self.monitor_registry[name] = gauge def counter_inc(self, name, label=None): @@ -141,8 +146,11 @@ def histogram_observe(self, name, value, label=None): else: self.monitor_registry[name].labels(model_name=self.model_name, method=label).observe(value) - def gauge_set(self, name, value): - self.monitor_registry[name].labels(model_name=self.model_name).set(value) + def gauge_set(self, name, value, labels=None): + metric_labels = {"model_name": self.model_name} + if labels: + metric_labels.update(labels) + self.monitor_registry[name].labels(**metric_labels).set(value) def push_metrices(self): if self.gateway_url is not None: