Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 49 additions & 11 deletions lightllm/server/httpserver_for_pd_master/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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", "")
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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()
Expand Down Expand Up @@ -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] = {}
Expand Down Expand Up @@ -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
Expand Down
5 changes: 3 additions & 2 deletions lightllm/server/metrics/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
16 changes: 12 additions & 4 deletions lightllm/server/metrics/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
}


Expand Down Expand Up @@ -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 [])
Expand All @@ -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):
Expand All @@ -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:
Expand Down
Loading