diff --git a/tensorrt_llm/commands/serve.py b/tensorrt_llm/commands/serve.py index c1013eb3c5cc..0c31ef07478a 100644 --- a/tensorrt_llm/commands/serve.py +++ b/tensorrt_llm/commands/serve.py @@ -3,7 +3,7 @@ import signal # Added import import subprocess # nosec B404 import sys -from typing import Any, List, Optional +from typing import Any, Optional import click import torch @@ -20,8 +20,7 @@ from tensorrt_llm.llmapi import (BuildConfig, CapacitySchedulerPolicy, DynamicBatchConfig, KvCacheConfig, SchedulerConfig) -from tensorrt_llm.llmapi.disagg_utils import (CtxGenServerConfig, - MetadataServerConfig, ServerRole, +from tensorrt_llm.llmapi.disagg_utils import (MetadataServerConfig, ServerRole, parse_disagg_config_file, parse_metadata_server_config_file) from tensorrt_llm.llmapi.llm_utils import update_llm_args_with_extra_dict @@ -442,19 +441,6 @@ def serve_encoder(model: str, host: str, port: int, log_level: str, launch_mm_encoder_server(host, port, encoder_args, metadata_server_cfg) -def get_ctx_gen_server_urls( - server_configs: List[CtxGenServerConfig]) -> List[str]: - ctx_server_urls = [] - gen_server_urls = [] - for cfg in server_configs: - if cfg.type == "ctx": - ctx_server_urls.append(f"http://{cfg.hostname}:{cfg.port}") - else: - gen_server_urls.append(f"http://{cfg.hostname}:{cfg.port}") - - return ctx_server_urls, gen_server_urls - - @click.command("disaggregated") @click.option("-c", "--config_file", @@ -491,22 +477,13 @@ def disaggregated(config_file: Optional[str], disagg_cfg = parse_disagg_config_file(config_file) - ctx_server_urls, gen_server_urls = get_ctx_gen_server_urls( - disagg_cfg.server_configs) - metadata_server_cfg = parse_metadata_server_config_file( metadata_server_config_file) - server = OpenAIDisaggServer( - ctx_servers=ctx_server_urls, - gen_servers=gen_server_urls, - req_timeout_secs=request_timeout, - server_start_timeout_secs=server_start_timeout, - max_retries=disagg_cfg.max_retries, - ctx_router_config=disagg_cfg.ctx_router_config, - gen_router_config=disagg_cfg.gen_router_config, - conditional_disagg_config=disagg_cfg.conditional_disagg_config, - metadata_server_cfg=metadata_server_cfg) + server = OpenAIDisaggServer(config=disagg_cfg, + req_timeout_secs=request_timeout, + server_start_timeout_secs=server_start_timeout, + metadata_server_cfg=metadata_server_cfg) asyncio.run(server(disagg_cfg.hostname, disagg_cfg.port)) diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index faa9bd08b34a..2f262aa30d19 100644 --- a/tensorrt_llm/llmapi/disagg_utils.py +++ b/tensorrt_llm/llmapi/disagg_utils.py @@ -52,6 +52,7 @@ class DisaggServerConfig(): gen_router_config: Optional[RouterConfig] = None conditional_disagg_config: Optional[ConditionalDisaggConfig] = None max_retries: int = 3 + perf_metrics_max_requests: int = 0 @dataclass @@ -63,6 +64,20 @@ class MetadataServerConfig(): refresh_interval: float = 10.0 +def get_ctx_gen_server_urls( + server_configs: list[CtxGenServerConfig] +) -> tuple[list[str], list[str]]: + ctx_server_urls = [] + gen_server_urls = [] + for cfg in server_configs: + if cfg.type == "ctx": + ctx_server_urls.append(f"http://{cfg.hostname}:{cfg.port}") + else: + gen_server_urls.append(f"http://{cfg.hostname}:{cfg.port}") + + return ctx_server_urls, gen_server_urls + + def parse_disagg_config_file(yaml_config_file: str): with open(yaml_config_file, 'r') as file: @@ -77,6 +92,7 @@ def parse_disagg_config_file(yaml_config_file: str): def extract_disagg_cfg(hostname: str = 'localhost', port: int = 8000, max_retries: int = 3, + perf_metrics_max_requests: int = 0, context_servers: Optional[dict] = None, generation_servers: Optional[dict] = None, conditional_disagg_config: Optional[dict] = None, @@ -115,7 +131,8 @@ def extract_disagg_cfg(hostname: str = 'localhost', config = DisaggServerConfig(server_configs, hostname, port, ctx_router_config, gen_router_config, - conditional_disagg_config, max_retries) + conditional_disagg_config, max_retries, + perf_metrics_max_requests) return config diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 6ed4dea76c7d..b714ea61e2fc 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -2115,6 +2115,12 @@ class TorchLlmArgs(BaseLlmArgs): description="Print iteration logs.", status="beta") + perf_metrics_max_requests: int = Field( + default=0, + description= + "The maximum number of requests for perf metrics. Must also set request_perf_metrics to true to get perf metrics.", + status="prototype") + batch_wait_timeout_ms: float = Field( default=0, description= diff --git a/tensorrt_llm/serve/openai_disagg_server.py b/tensorrt_llm/serve/openai_disagg_server.py index 85a052636ba4..4a9b409be925 100644 --- a/tensorrt_llm/serve/openai_disagg_server.py +++ b/tensorrt_llm/serve/openai_disagg_server.py @@ -1,12 +1,14 @@ #!/usr/bin/env python import asyncio import copy +import itertools import os import signal import traceback +from collections import deque from contextlib import asynccontextmanager from http import HTTPStatus -from typing import List, Optional, Type, Union +from typing import Optional, Type, Union import aiohttp import uvicorn @@ -17,9 +19,9 @@ # yapf: disable from tensorrt_llm.executor import CppExecutorError -from tensorrt_llm.llmapi.disagg_utils import (ConditionalDisaggConfig, +from tensorrt_llm.llmapi.disagg_utils import (DisaggServerConfig, MetadataServerConfig, - RouterConfig) + get_ctx_gen_server_urls) from tensorrt_llm.logger import logger from tensorrt_llm.serve.metadata_server import create_metadata_server from tensorrt_llm.serve.openai_protocol import (ChatCompletionRequest, @@ -37,35 +39,45 @@ class OpenAIDisaggServer: def __init__(self, - ctx_servers: List[str], - gen_servers: List[str], + config: DisaggServerConfig, req_timeout_secs: int = 180, server_start_timeout_secs: int = 180, - max_retries: int = 3, - ctx_router_config: Optional[RouterConfig] = None, - gen_router_config: Optional[RouterConfig] = None, - conditional_disagg_config: Optional[ConditionalDisaggConfig] = None, metadata_server_cfg: Optional[MetadataServerConfig] = None): - self.ctx_servers = ctx_servers - self.gen_servers = gen_servers + self.ctx_servers, self.gen_servers = get_ctx_gen_server_urls(config.server_configs) self.metadata_server = create_metadata_server(metadata_server_cfg) - self.ctx_router = create_router(ctx_router_config, ctx_servers, metadata_server_cfg, self.metadata_server) - self.gen_router = create_router(gen_router_config, gen_servers, metadata_server_cfg, self.metadata_server) - self.conditional_disagg_config = conditional_disagg_config + self.ctx_router = create_router( + config.ctx_router_config, self.ctx_servers, metadata_server_cfg, self.metadata_server) + self.gen_router = create_router( + config.gen_router_config, self.gen_servers, metadata_server_cfg, self.metadata_server) + self.conditional_disagg_config = config.conditional_disagg_config + + self.perf_metrics_max_requests = config.perf_metrics_max_requests + if self.perf_metrics_max_requests > 0: + # record corresponding keys of context and generation servers for perf metrics + # (ctx_server, gen_server, ctx_request_id) + self.perf_metrics_keys = deque(maxlen=self.perf_metrics_max_requests) + self.perf_metrics_keys_lock = asyncio.Lock() + # server_key -> {ctx_request_id: perf_metrics} + self.server_perf_metrics: dict[str, dict[int, dict]] = {} + else: + self.perf_metrics_keys = None + self.perf_metrics_keys_lock = None + self.server_perf_metrics = None - if max_retries < 0: - raise ValueError(f"Max retries {max_retries} must be greater than or equal to 0") - self.max_retries = max_retries + if config.max_retries < 0: + raise ValueError(f"Max retries {config.max_retries} must be greater than or equal to 0") + self.max_retries = config.max_retries logger.info(f"Server max retries: {self.max_retries}") if (len(self.gen_servers) == 0): raise ValueError("At least one generation server must be provided") - if os.getenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY") != "1" and len(ctx_servers) == 0: + if os.getenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY") != "1" and len(self.ctx_servers) == 0: raise ValueError("At least one context server must be provided") - if self.conditional_disagg_config is not None and not isinstance(self.gen_router, KvCacheAwareRouter): + if self.conditional_disagg_config is not None and \ + not isinstance(self.gen_router, KvCacheAwareRouter): raise ValueError("Generation router must be a KvCacheAwareRouter to enable conditional disaggregation") # Session will be initialized in lifespan @@ -112,6 +124,7 @@ def create_error_response( def register_routes(self): self.app.add_api_route("/health", self.health, methods=["GET"]) self.app.add_api_route("/version", self.version, methods=["GET"]) + self.app.add_api_route("/perf_metrics", self.perf_metrics, methods=["GET"]) self.app.add_api_route("/v1/completions", self.openai_completion, methods=["POST"]) @@ -126,6 +139,61 @@ async def version(self) -> JSONResponse: ver = {"version": VERSION} return JSONResponse(content=ver) + async def _add_perf_metrics_keys(self, ctx_server: str, gen_server: str, ctx_request_id: int): + async with self.perf_metrics_keys_lock: + self.perf_metrics_keys.append((ctx_server, gen_server, ctx_request_id)) + + async def perf_metrics(self) -> JSONResponse: + if self.perf_metrics_keys is None: + return JSONResponse(content=[]) + + perf_metrics = {} + exc = None + try: + for server in self.ctx_servers + self.gen_servers: + async with self.session.get(f"{server}/perf_metrics") as response: + server_perf_metrics = await response.json() + perf_metrics[server] = server_perf_metrics + except Exception as e: + # Keep the exception to raise it after saving perf metrics + exc = e + + return_metrics = [] + async with self.perf_metrics_keys_lock: + for server in perf_metrics: + server_metrics = self.server_perf_metrics.setdefault(server, {}) + for request_perf_metrics in perf_metrics[server]: + ctx_request_id = request_perf_metrics.get("ctx_request_id", None) + if ctx_request_id is None: + continue + server_metrics[ctx_request_id] = request_perf_metrics + + if len(server_metrics) > self.perf_metrics_max_requests: + # Remove oldest requests and keep at most perf_metrics_max_requests + num_remove = len(server_metrics) - self.perf_metrics_max_requests + removed_keys = list(itertools.islice(server_metrics.keys(), num_remove)) + for ctx_request_id in removed_keys: + server_metrics.pop(ctx_request_id) + if exc is not None: + raise exc + + remain_keys = [] + for ctx_server, gen_server, ctx_request_id in self.perf_metrics_keys: + gen_perf_metrics = self.server_perf_metrics[gen_server].pop(ctx_request_id, None) + if gen_perf_metrics is None: + # generation not finished + remain_keys.append((ctx_server, gen_server, ctx_request_id)) + continue + ctx_perf_metrics = self.server_perf_metrics[ctx_server].pop(ctx_request_id, None) + return_metrics.append({ + "ctx_server": ctx_server, + "gen_server": gen_server, + "ctx_perf_metrics": ctx_perf_metrics, + "gen_perf_metrics": gen_perf_metrics}) + self.perf_metrics_keys = deque(remain_keys, maxlen=self.perf_metrics_max_requests) + + return JSONResponse(content=return_metrics) + async def merge_streaming_responses(self, ctx_response, gen_server: str, gen_req: Union[CompletionRequest, ChatCompletionRequest]): @@ -263,6 +331,10 @@ async def _send_disagg_request(self, req: Union[CompletionRequest, ChatCompletio gen_server, _ = await self.gen_router.get_next_server(req) logger.debug("Sending request to gen server: %s", gen_server) + if need_ctx and self.perf_metrics_keys is not None: + asyncio.create_task(self._add_perf_metrics_keys( + ctx_server, gen_server, req.disaggregated_params.ctx_request_id)) + if not req.stream: try: #If request finished after first token for reason other than length, return right away and skip gen diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index cecbdb5d1d4c..54d0733f7761 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -4,6 +4,7 @@ import re import signal import traceback +from collections import deque from contextlib import asynccontextmanager from datetime import datetime from http import HTTPStatus @@ -84,12 +85,18 @@ def __init__(self, else: self.model = model self.metrics_collector = None + self.perf_metrics = None + self.perf_metrics_lock = None if self.llm.args.return_perf_metrics: set_prometheus_multiproc_dir() self.metrics_collector = MetricsCollector({ "model_name": "undefined", "engine_type": "undefined" }) + max_perf_metrics = self.llm.args.perf_metrics_max_requests + if max_perf_metrics > 0: + self.perf_metrics = deque(maxlen=max_perf_metrics) + self.perf_metrics_lock = asyncio.Lock() @asynccontextmanager async def lifespan(app: FastAPI): @@ -159,6 +166,7 @@ def register_routes(self): self.app.add_api_route("/v1/models", self.get_model, methods=["GET"]) # TODO: the metrics endpoint only reports iteration stats, not the runtime stats for now self.app.add_api_route("/metrics", self.get_iteration_stats, methods=["GET"]) + self.app.add_api_route("/perf_metrics", self.get_perf_metrics, methods=["GET"]) # TODO: workaround before ETCD support self.app.add_api_route("/kv_cache_events", self.get_kv_cache_events, methods=["POST"]) self.app.add_api_route("/v1/completions", @@ -255,6 +263,50 @@ async def get_iteration_stats(self) -> JSONResponse: stats.append(stat) return JSONResponse(content=stats) + async def get_perf_metrics(self) -> JSONResponse: + if self.perf_metrics is None: + return JSONResponse(content=[]) + async with self.perf_metrics_lock: + perf_metrics = self.perf_metrics + self.perf_metrics = deque(maxlen=self.llm.args.perf_metrics_max_requests) + for metrics_dict in perf_metrics: + metrics = metrics_dict["perf_metrics"] + timing_metrics = metrics.timing_metrics + kv_cache_metrics = metrics.kv_cache_metrics + speculative_decoding = metrics.speculative_decoding + metrics_json = { + "first_iter": metrics.first_iter, + "last_iter": metrics.last_iter, + # exclude metrics.iter since it is only meaningful when the request is not finished + } + metrics_json["timing_metrics"] = { + "arrival_time": timing_metrics.arrival_time.total_seconds(), + "first_scheduled_time": timing_metrics.first_scheduled_time.total_seconds(), + "first_token_time": timing_metrics.first_token_time.total_seconds(), + "last_token_time": timing_metrics.last_token_time.total_seconds(), + } + metrics_json["kv_cache_metrics"] = { + "num_total_allocated_blocks": kv_cache_metrics.num_total_allocated_blocks, + "num_new_allocated_blocks": kv_cache_metrics.num_new_allocated_blocks, + "num_reused_blocks": kv_cache_metrics.num_reused_blocks, + "num_missed_blocks": kv_cache_metrics.num_missed_blocks, + } + if timing_metrics.kv_cache_size > 0: + metrics_json["timing_metrics"].update({ + # TODO: move to kv_cache_metrics + "kv_cache_size": timing_metrics.kv_cache_size, + "kv_cache_transfer_start": timing_metrics.kv_cache_transfer_start.total_seconds(), + "kv_cache_transfer_end": timing_metrics.kv_cache_transfer_end.total_seconds(), + }) + if speculative_decoding.total_draft_tokens > 0: + metrics_json["speculative_decoding"] = { + "acceptance_rate": speculative_decoding.acceptance_rate, + "total_accepted_draft_tokens": speculative_decoding.total_accepted_draft_tokens, + "total_draft_tokens": speculative_decoding.total_draft_tokens, + } + metrics_dict["perf_metrics"] = metrics_json + return JSONResponse(content=list(perf_metrics)) + async def get_kv_cache_events(self) -> JSONResponse: events = [] try: @@ -265,6 +317,23 @@ async def get_kv_cache_events(self) -> JSONResponse: pass return JSONResponse(content=events) + async def _extract_metrics(self, res: RequestOutput): + if not res.finished: + return + if self.metrics_collector: + self.metrics_collector.log_metrics_dict(res.metrics_dict) + if self.llm.args.return_perf_metrics: + output = res.outputs[0] + item = { + "request_id": res.request_id, + "perf_metrics": res.outputs[0].request_perf_metrics + } + if output.disaggregated_params: + item["ctx_request_id"] = output.disaggregated_params.ctx_request_id + if self.perf_metrics is not None: + async with self.perf_metrics_lock: + self.perf_metrics.append(item) + async def openai_chat(self, request: ChatCompletionRequest, raw_request: Request) -> Response: def get_role() -> str: @@ -280,8 +349,7 @@ async def chat_stream_generator( post_processor, args = postproc_params.post_processor, postproc_params.postproc_args async for res in promise: pp_results = res.outputs[0]._postprocess_result if self.postproc_worker_enabled else post_processor(res, args) - if res.finished and self.metrics_collector: - self.metrics_collector.log_metrics_dict(res.metrics_dict) + await self._extract_metrics(res) for pp_res in pp_results: yield pp_res yield "data: [DONE]\n\n" @@ -299,8 +367,7 @@ async def create_chat_response( # Add prompt_tokens_ids to the response if disaggregated_params and disaggregated_params.request_type and disaggregated_params.request_type == "context_only": chat_response.prompt_token_ids = promise.prompt_token_ids - if promise.finished and self.metrics_collector: - self.metrics_collector.log_metrics_dict(promise.metrics_dict) + await self._extract_metrics(promise) return chat_response try: @@ -469,8 +536,7 @@ async def completion_response(promise: RequestOutput, if disaggregated_params and disaggregated_params.request_type and disaggregated_params.request_type == "context_only": # Include prompt token ids for context-only requests pp_result.prompt_token_ids = response.prompt_token_ids - if response.finished and self.metrics_collector: - self.metrics_collector.log_metrics_dict(response.metrics_dict) + await self._extract_metrics(response) return pp_result def merge_completion_responses(responses: List[CompletionResponse]) -> CompletionResponse: @@ -506,8 +572,7 @@ async def completion_generator(promise: RequestOutput, params: Optional[Postproc pp_result = post_processor(output, args) else: pp_result = output.outputs[0]._postprocess_result - if output.finished and self.metrics_collector: - self.metrics_collector.log_metrics_dict(output.metrics_dict) + await self._extract_metrics(output) for pp_res in pp_result: yield pp_res diff --git a/tests/integration/defs/disaggregated/test_configs/disagg_config_metrics.yaml b/tests/integration/defs/disaggregated/test_configs/disagg_config_metrics.yaml new file mode 100644 index 000000000000..6d566aa4f99b --- /dev/null +++ b/tests/integration/defs/disaggregated/test_configs/disagg_config_metrics.yaml @@ -0,0 +1,28 @@ +hostname: localhost +port: 8000 +model: TinyLlama/TinyLlama-1.1B-Chat-v1.0 +free_gpu_memory_fraction: 0.25 +backend: "pytorch" +cuda_graph_config: null +disable_overlap_scheduler: True +perf_metrics_max_requests: 1000 +context_servers: + num_instances: 1 + tensor_parallel_size: 1 + pipeline_parallel_size: 1 + return_perf_metrics: True + perf_metrics_max_requests: 1000 + cache_transceiver_config: + backend: DEFAULT + urls: + - "localhost:8001" +generation_servers: + num_instances: 1 + tensor_parallel_size: 1 + pipeline_parallel_size: 1 + return_perf_metrics: True + perf_metrics_max_requests: 1000 + cache_transceiver_config: + backend: DEFAULT + urls: + - "localhost:8002" diff --git a/tests/integration/defs/disaggregated/test_disaggregated.py b/tests/integration/defs/disaggregated/test_disaggregated.py index 24000b1f802b..3b8e54fe6a0e 100644 --- a/tests/integration/defs/disaggregated/test_disaggregated.py +++ b/tests/integration/defs/disaggregated/test_disaggregated.py @@ -17,6 +17,7 @@ import re import subprocess import tempfile +from typing import Callable import pytest import yaml @@ -35,6 +36,14 @@ def cleanup_output_files(): pass +def get_disagg_server_url_from_cfg(config_file: str) -> str: + with open(config_file, 'r') as file: + config = yaml.safe_load(file) + server_host = config.get('hostname', 'localhost') + server_port = config.get('port', 8000) + return f"http://{server_host}:{server_port}" + + def get_test_config(test_desc, example_dir, test_root): """Get test configuration based on test description.""" test_configs_root = f"{test_root}/test_configs" @@ -57,6 +66,7 @@ def get_test_config(test_desc, example_dir, test_root): (2, f"{test_configs_root}/disagg_config_cuda_graph_padding.yaml"), "mixed": (2, f"{test_configs_root}/disagg_config_mixed.yaml"), "overlap": (2, f"{test_configs_root}/disagg_config_overlap.yaml"), + "perf_metrics": (2, f"{test_configs_root}/disagg_config_metrics.yaml"), "trtllm_sampler": (2, f"{test_configs_root}/disagg_config_trtllm_sampler.yaml"), "load_balance": @@ -152,7 +162,8 @@ def run_disaggregated_test(example_dir, num_iters=5, env=None, cwd=None, - prompt_file="prompts.json"): + prompt_file="prompts.json", + extra_endpoints_test: Callable[[str], None] = None): """Run disaggregated test with given configuration.""" cleanup_output_files() run_env = env.copy() @@ -172,6 +183,7 @@ def run_disaggregated_test(example_dir, 'trtllm-serve', 'disaggregated', '--server_start_timeout', str(server_start_timeout), '-c', config_file ] + server_url = get_disagg_server_url_from_cfg(config_file) try: with ( # Start workers @@ -232,6 +244,9 @@ def run_disaggregated_test(example_dir, if prompt_file == "long_prompts.json": continue + if extra_endpoints_test is not None: + extra_endpoints_test(server_url) + # Verify outputs not_expected_strings = ["Berlin Berlin"] @@ -511,6 +526,57 @@ def test_disaggregated_overlap(disaggregated_test_root, llm_venv, cwd=llm_venv.get_working_directory()) +@pytest.mark.parametrize("llama_model_root", ['TinyLlama-1.1B-Chat-v1.0'], + indirect=True) +def test_disaggregated_perf_metrics(disaggregated_test_root, llm_venv, + disaggregated_example_root, + llama_model_root): + src_dst_dict = { + llama_model_root: + f"{llm_venv.get_working_directory()}/TinyLlama/TinyLlama-1.1B-Chat-v1.0", + } + for src, dst in src_dst_dict.items(): + if not os.path.islink(dst): + os.makedirs(os.path.dirname(dst), exist_ok=True) + os.symlink(src, dst, target_is_directory=True) + + def extra_endpoints_test(server_url: str): + import json + import urllib.request + + with urllib.request.urlopen(f"{server_url}/perf_metrics", + timeout=10) as resp: + assert resp.status == 200 + perf_metrics = json.load(resp) + assert len(perf_metrics) > 0 + item = perf_metrics[0] + assert "ctx_server" in item + assert "gen_server" in item + assert "ctx_perf_metrics" in item + assert "gen_perf_metrics" in item + assert item["ctx_perf_metrics"]["ctx_request_id"] == item[ + "gen_perf_metrics"]["ctx_request_id"] + ctx_metrics = item["ctx_perf_metrics"]["perf_metrics"]["timing_metrics"] + gen_metrics = item["gen_perf_metrics"]["perf_metrics"]["timing_metrics"] + # only one token is generated in ctx + assert ctx_metrics["last_token_time"] - ctx_metrics[ + "first_token_time"] < 1e-3 + assert ctx_metrics["last_token_time"] < gen_metrics["arrival_time"] + assert gen_metrics["kv_cache_size"] > 0 + assert gen_metrics["arrival_time"] < gen_metrics[ + "kv_cache_transfer_start"] + assert gen_metrics["kv_cache_transfer_start"] < gen_metrics[ + "kv_cache_transfer_end"] + assert gen_metrics["kv_cache_transfer_end"] < gen_metrics[ + "first_scheduled_time"] + + run_disaggregated_test(disaggregated_example_root, + "perf_metrics", + env=llm_venv._new_env, + cwd=llm_venv.get_working_directory(), + extra_endpoints_test=extra_endpoints_test) + + @pytest.mark.parametrize("llama_model_root", ['TinyLlama-1.1B-Chat-v1.0'], indirect=True) def test_disaggregated_trtllm_sampler(disaggregated_test_root, llm_venv, diff --git a/tests/integration/defs/test_e2e.py b/tests/integration/defs/test_e2e.py index ef615843bbc2..9a1223cb53a3 100644 --- a/tests/integration/defs/test_e2e.py +++ b/tests/integration/defs/test_e2e.py @@ -1499,6 +1499,13 @@ def test_openai_chat_with_logit_bias(llm_root, llm_venv, sampler: str): ]) +def test_openai_perf_metrics(llm_root, llm_venv): + test_root = unittest_path() / "llmapi" / "apps" + llm_venv.run_cmd( + ["-m", "pytest", + str(test_root / "_test_openai_perf_metrics.py")]) + + def test_openai_prometheus(llm_root, llm_venv): test_root = unittest_path() / "llmapi" / "apps" llm_venv.run_cmd( diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index 30fc6c05b58c..df8852d78598 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -25,6 +25,7 @@ l0_a10: - test_e2e.py::test_openai_chat_structural_tag_example - test_e2e.py::test_openai_chat_json_example - test_e2e.py::test_openai_chat_multimodal_example + - test_e2e.py::test_openai_perf_metrics - test_e2e.py::test_openai_prometheus - test_e2e.py::test_openai_lora - test_e2e.py::test_trtllm_serve_multimodal_example diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 50575e730bb7..68b27614db93 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -73,6 +73,7 @@ l0_h100: - disaggregated/test_disaggregated.py::test_disaggregated_cuda_graph[TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated.py::test_disaggregated_mixed[TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated.py::test_disaggregated_overlap[TinyLlama-1.1B-Chat-v1.0] + - disaggregated/test_disaggregated.py::test_disaggregated_perf_metrics[TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated_single_gpu.py::test_disaggregated_simple_llama[False-False-TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated_single_gpu.py::test_disaggregated_simple_llama[False-True-TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated_single_gpu.py::test_disaggregated_simple_llama[True-False-TinyLlama-1.1B-Chat-v1.0] diff --git a/tests/unittest/api_stability/references/llm.yaml b/tests/unittest/api_stability/references/llm.yaml index d9dcd0f83d2f..7e1958ff6452 100644 --- a/tests/unittest/api_stability/references/llm.yaml +++ b/tests/unittest/api_stability/references/llm.yaml @@ -131,6 +131,10 @@ methods: annotation: bool default: False status: beta + perf_metrics_max_requests: + annotation: int + default: 0 + status: prototype torch_compile_config: annotation: Optional[tensorrt_llm.llmapi.llm_args.TorchCompileConfig] default: null diff --git a/tests/unittest/llmapi/apps/_test_openai_perf_metrics.py b/tests/unittest/llmapi/apps/_test_openai_perf_metrics.py new file mode 100644 index 000000000000..c1049939953b --- /dev/null +++ b/tests/unittest/llmapi/apps/_test_openai_perf_metrics.py @@ -0,0 +1,101 @@ +import json +import logging +import os +import tempfile +from urllib.request import urlopen + +import pytest +import yaml + +from ..test_llm import get_model_path +from .openai_server import RemoteOpenAIServer + +# Configure logging +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + + +@pytest.fixture(scope="module", ids=["TinyLlama-1.1B-Chat"]) +def model_name(): + return "llama-models-v2/TinyLlama-1.1B-Chat-v1.0" + + +@pytest.fixture(scope="module") +def temp_extra_llm_api_options_file(request): + temp_dir = tempfile.gettempdir() + temp_file_path = os.path.join(temp_dir, "extra_llm_api_options.yaml") + try: + extra_llm_api_options_dict = { + "return_perf_metrics": True, + "perf_metrics_max_requests": 10 + } + + with open(temp_file_path, 'w') as f: + yaml.dump(extra_llm_api_options_dict, f) + + yield temp_file_path + finally: + if os.path.exists(temp_file_path): + os.remove(temp_file_path) + + +@pytest.fixture(scope="module") +def server(model_name: str, + temp_extra_llm_api_options_file: str) -> RemoteOpenAIServer: + model_path = get_model_path(model_name) + args = ["--backend", "pytorch", "--tp_size", "1"] + args.extend(["--extra_llm_api_options", temp_extra_llm_api_options_file]) + logger.info(f"Starting server, model: {model_name}, args: {args}") + with RemoteOpenAIServer(model_path, args) as remote_server: + yield remote_server + logger.info("Tests completed, shutting down server") + + +def test_metrics_endpoint(server: RemoteOpenAIServer): + + client = server.get_client() + client.completions.create( + model="Server", + prompt="Hello, my name is", + max_tokens=25, + stream=False, + ) + + response = urlopen(f'{server.url_root}/perf_metrics') + assert response.status is 200 + + data_list = json.loads(response.read()) + assert len(data_list) == 1 + assert "perf_metrics" in data_list[0] + assert "request_id" in data_list[0] + + data = data_list[0]["perf_metrics"] + assert "first_iter" in data + assert "last_iter" in data + assert data["first_iter"] <= data["last_iter"] + + timing_metrics = data["timing_metrics"] + assert "arrival_time" in timing_metrics + assert "first_scheduled_time" in timing_metrics + assert "first_token_time" in timing_metrics + assert "last_token_time" in timing_metrics + assert timing_metrics["arrival_time"] < timing_metrics[ + "first_scheduled_time"] + assert timing_metrics["first_scheduled_time"] < timing_metrics[ + "first_token_time"] + assert timing_metrics["first_token_time"] <= timing_metrics[ + "last_token_time"] + + kv_cache_metrics = data["kv_cache_metrics"] + assert "num_total_allocated_blocks" in kv_cache_metrics + assert "num_new_allocated_blocks" in kv_cache_metrics + assert "num_reused_blocks" in kv_cache_metrics + assert "num_missed_blocks" in kv_cache_metrics + assert kv_cache_metrics["num_new_allocated_blocks"] <= kv_cache_metrics[ + "num_total_allocated_blocks"] + + # exclude disagg specific metrics + assert "ctx_request_id" not in data_list[0] + assert "kv_cache_size" not in timing_metrics + assert "kv_cache_transfer_start" not in timing_metrics + assert "kv_cache_transfer_end" not in timing_metrics